From a0985b1b3b760fce6339e36b08ccb1c37e9076eb Mon Sep 17 00:00:00 2001 From: yanqiangmiffy <1185918903@qq.com> Date: Tue, 9 Dec 2025 16:53:58 +0800 Subject: [PATCH 1/4] update@merge pr --- docs/git.md | 75 +++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 75 insertions(+) diff --git a/docs/git.md b/docs/git.md index c976c29..99dc236 100644 --- a/docs/git.md +++ b/docs/git.md @@ -11,3 +11,78 @@ git merge main git reset --hard origin/main ``` +## 合并PR + +这是一个非常经典且常见的 Git 协作场景。你现在的目标是将 `main` 分支上(包含那两个修复安装问题的 PR)的代码,同步到你当前正在工作的 `pipeline` 分支上。 + +最推荐、最稳妥的方法是使用 **Merge(合并)**。 + +### 方法一:标准合并 (Recommended) + +这种方法最安全,它会把 `main` 上的新内容“合”过来,并保留详细的历史记录。 + +请在你的终端(当前已经在 `pipeline` 分支)按顺序执行以下步骤: + +**1. 切换到 main 分支并更新代码** +首先,你需要确保你本地的 `main` 分支是最新的(包含了你在网页端合并的那两个 PR)。 + +```bash +git checkout main +git pull origin main +``` + +**2. 切换回 pipeline 分支** +更新完 `main` 后,回到你的工作分支。 + +```bash +git checkout pipeline +``` + +**3. 将 main 合并进 pipeline** +这一步会将 `main` 的修补内容同步到你当前的 `pipeline` 分支。 + +```bash +git merge main +``` + +[Image of git merge diagram showing feature branch updating from main] + +**4. 处理可能出现的冲突 (Conflict)** + + * **如果一切顺利**:终端会提示 "Merge made by the 'ort' strategy" 或类似信息,你不需要做任何额外操作。 + * **如果有冲突**:Git 会提示 `CONFLICT (content)...`。 + * 你需要打开冲突的文件,找到 `<<<<<<<` 和 `>>>>>>>` 标记的地方,手动修改代码保留你需要的部分。 + * 修改完后,执行 `git add .` 和 `git commit` 来完成合并。 + +----- + +### 方法二:变基 (Rebase) —— 高级选项 + +如果你希望你的 `pipeline` 分支历史非常干净(看起来就像是你是在那两个 PR 修复**之后**才开始开发 `pipeline` 的),你可以使用 Rebase。 + +**注意:** 如果 `pipeline` 分支是多人协作的(即别人也在这个分支上写代码),**不要**用这个方法。如果是你自己单独开发的,这个方法更优雅。 + +```bash +# 1. 更新 main (同上) +git checkout main +git pull origin main + +# 2. 回到 pipeline +git checkout pipeline + +# 3. 执行变基 +git rebase main +``` + + * **它的作用**:它会把你在 `pipeline` 上写的代码先“拿下来”,把 `main` 的新代码放进去,然后再把你写的代码“贴”在最后面。这样你的提交历史是一条直线。 + +----- + +### 总结 + +对于你的情况,**我建议直接使用方法一(Merge)**。 + +虽然 Rebase 历史好看,但在处理冲突时可能会稍微复杂一点(每一个 commit 都可能冲突)。Merge 是最直观的方式,能立刻让你拥有 `main` 上的修复补丁,继续开发。 + +**现在的下一步:** +你可以直接在终端试着运行 **方法一** 的命令。如果在第 3 步(merge)遇到了冲突(Conflict)不知道怎么修,你可以把报错信息发给我,我来教你如何解决。 \ No newline at end of file From d8001d6bde5e8dc63ae7acf44469e20a663b23f6 Mon Sep 17 00:00:00 2001 From: yanqiangmiffy <1185918903@qq.com> Date: Tue, 9 Dec 2025 18:19:31 +0800 Subject: [PATCH 2/4] bug@fixed chunk --- trustrag/applications/rag_openai.py | 6 ++++++ trustrag/modules/document/chunk.py | 32 +++++++++++++++++++++++++++-- 2 files changed, 36 insertions(+), 2 deletions(-) diff --git a/trustrag/applications/rag_openai.py b/trustrag/applications/rag_openai.py index c925bd5..5fd2e24 100644 --- a/trustrag/applications/rag_openai.py +++ b/trustrag/applications/rag_openai.py @@ -67,6 +67,12 @@ def init_vector_store(self): pass print("chunking for paragraphs") for paragraphs in all_paragraphs: + if paragraphs: + first_para = paragraphs[0] + if isinstance(first_para, dict): + print(first_para.keys()) + else: + print(f"paragraph type: {type(first_para)}") chunks = self.tc.get_chunks(paragraphs, 256) all_chunks.extend(chunks) self.retriever.build_from_texts(all_chunks) diff --git a/trustrag/modules/document/chunk.py b/trustrag/modules/document/chunk.py index 4355666..6c7826a 100644 --- a/trustrag/modules/document/chunk.py +++ b/trustrag/modules/document/chunk.py @@ -107,6 +107,31 @@ def split_large_sentence(self, sentence: str, chunk_size: int) -> list[str]: return sentence_parts + def _normalize_paragraphs(self, paragraphs: list) -> list[str]: + """ + Normalize paragraphs into a list of strings. + + Supports paragraphs passed as strings or dictionaries (e.g. {"title": ..., "content": ...}). + Falls back to str() for any other types to avoid runtime errors. + """ + normalized = [] + for p in paragraphs: + if isinstance(p, str): + normalized.append(p) + elif isinstance(p, dict): + # Prefer common text-bearing fields; join with newlines to preserve separation. + parts = [] + for key in ("title", "content", "text", "body"): + if key in p and p[key]: + parts.append(str(p[key])) + if parts: + normalized.append("\n".join(parts)) + else: + normalized.append(str(p)) + else: + normalized.append(str(p)) + return normalized + def get_chunks(self, paragraphs: list[str], chunk_size: int) -> list[str]: """ Splits a list of paragraphs into chunks based on a specified token size. @@ -118,15 +143,18 @@ def get_chunks(self, paragraphs: list[str], chunk_size: int) -> list[str]: Returns: list[str]: A list of text chunks, each containing sentences that fit within the token limit. """ + # Normalize paragraphs to string list to handle dict inputs gracefully + normalized_paragraphs = self._normalize_paragraphs(paragraphs) + # Combine paragraphs into a single text - text = ''.join(paragraphs) + text = ''.join(normalized_paragraphs) # Split the text into sentences sentences = self.split_sentences(text) # If no sentences are found, treat paragraphs as sentences if len(sentences) == 0: - sentences = paragraphs + sentences = normalized_paragraphs chunks = [] current_chunk = [] From b1b6e52bc631ddec14c0bc92e917103e8a0f0b6f Mon Sep 17 00:00:00 2001 From: yanqiangmiffy <1185918903@qq.com> Date: Wed, 10 Dec 2025 17:23:20 +0800 Subject: [PATCH 3/4] update@api fixed --- api/rag/apps/core/judge/views.py | 7 ++-- api/rag/apps/core/rerank/views.py | 7 +--- examples/generator/openai_chat_example.py | 47 ++++++++++++++++++++--- 3 files changed, 47 insertions(+), 14 deletions(-) diff --git a/api/rag/apps/core/judge/views.py b/api/rag/apps/core/judge/views.py index bcda596..37be325 100644 --- a/api/rag/apps/core/judge/views.py +++ b/api/rag/apps/core/judge/views.py @@ -11,16 +11,17 @@ """ import loguru from fastapi import APIRouter -from trustrag.config.config_loader import config +from trustrag.config.config_loader import ConfigLoader from api.rag.apps.core.judge.bodys import JudgeBody from api.rag.apps.handle.response.json_response import ApiResponse from trustrag.modules.judger.bge_judger import BgeJudger, BgeJudgerConfig from trustrag.modules.judger.chatgpt_judger import OpenaiJudger, OpenaiJudgerConfig judge_router = APIRouter() - +# 使用本地配置文件初始化配置加载器(单例) +config = ConfigLoader(config_path="config_local.json") # 加载服务和模型配置 -llm_service = config.get_config('services.dmx') +llm_service = config.get_config('services.gomall') llm_model = config.get_config('models.llm') rerank_model = config.get_config('models.reranker') diff --git a/api/rag/apps/core/rerank/views.py b/api/rag/apps/core/rerank/views.py index c72f2dd..e235b3b 100644 --- a/api/rag/apps/core/rerank/views.py +++ b/api/rag/apps/core/rerank/views.py @@ -16,12 +16,9 @@ from api.rag.apps.core.rerank.models import Application from api.rag.apps.handle.response.json_response import UserNotFoundResponse, ApiResponse from trustrag.modules.reranker.bge_reranker import BgeReranker, BgeRerankerConfig -from trustrag.config.config_loader import config - -# from apps.handle.exception.exception import MallException -# from apps.core.config.models import LLMModel -# from tortoise.contrib.pydantic import pydantic_model_creator +from trustrag.config.config_loader import ConfigLoader +config = ConfigLoader(config_path="config_local.json") rerank_router = APIRouter() # 从配置文件加载重排序配置 rerank_service = config.get_config('services.rerank') diff --git a/examples/generator/openai_chat_example.py b/examples/generator/openai_chat_example.py index 2962370..f16cf58 100644 --- a/examples/generator/openai_chat_example.py +++ b/examples/generator/openai_chat_example.py @@ -1,3 +1,38 @@ +# from openai import OpenAI +# import os +# from dotenv import load_dotenv +# load_dotenv() +# +# # for key, value in os.environ.items(): +# # print(f"{key} = {value}") +# client = OpenAI( +# # 替换为您需要调用的模型服务Base Url +# base_url=os.environ.get("VOLCENGINE_BASE_URL"), +# # 环境变量中配置您的API Key +# api_key=os.environ.get("VOLCENGINE_API_KEY") +# ) +# +# +# print("----- standard request -----") +# completion = client.chat.completions.create( +# model="deepseek-r1-250120", +# messages = [ +# # {"role": "system", "content": "你是豆包,是由字节跳动开发的 AI 人工智能助手"}, +# # {"role": "user", "content": "常见的十字花科植物有哪些?"}, +# +# {"role": "system", "content": "你是豆包,是由字节跳动开发的 AI 人工智能助手"}, +# {"role": "user", "content": "请帮我生成一个json数据,不要输出额外内容,保证json能够正确解析?"}, +# ], +# ) +# print(completion) +# print("reasoning_content:\n",completion.choices[0].message.reasoning_content) +# print("------"*100) +# +# print("content:\n") +# print(completion.choices[0].message.content) +# + + from openai import OpenAI import os from dotenv import load_dotenv @@ -7,15 +42,15 @@ # print(f"{key} = {value}") client = OpenAI( # 替换为您需要调用的模型服务Base Url - base_url=os.environ.get("VOLCENGINE_BASE_URL"), + base_url="http://10.208.61.1:32004/api/v1/1504_gomall_qwen3/Qwen3-30B-A3B-Instruct-2507/", # 环境变量中配置您的API Key - api_key=os.environ.get("VOLCENGINE_API_KEY") + api_key="" ) print("----- standard request -----") completion = client.chat.completions.create( - model="deepseek-r1-250120", + model="Qwen3-30B-A3B-Instruct-2507", messages = [ # {"role": "system", "content": "你是豆包,是由字节跳动开发的 AI 人工智能助手"}, # {"role": "user", "content": "常见的十字花科植物有哪些?"}, @@ -24,9 +59,6 @@ {"role": "user", "content": "请帮我生成一个json数据,不要输出额外内容,保证json能够正确解析?"}, ], ) -print(completion) -print("reasoning_content:\n",completion.choices[0].message.reasoning_content) -print("------"*100) print("content:\n") print(completion.choices[0].message.content) @@ -34,3 +66,6 @@ + + + From 60cce2163d3d6f636b0325f56a42c98725a1ecde Mon Sep 17 00:00:00 2001 From: yanqiangmiffy <1185918903@qq.com> Date: Thu, 11 Dec 2025 17:39:44 +0800 Subject: [PATCH 4/4] update@fixed api --- Dockerfile | 2 +- api/README.md | 24 +++++ api/rag/.env.example | 32 +++++++ api/rag/apps/config/__init__.py | 14 +-- api/rag/apps/config/app_config.py | 23 ++--- api/rag/apps/config/base_config.py | 40 +++++++++ api/rag/apps/config/rerank_config.py | 18 ++-- api/rag/apps/core/judge/views.py | 17 ++-- api/rag/apps/core/rerank/views.py | 12 +-- api/rag/apps/core/rewrite/views.py | 7 +- api/rag/citation.json | 13 +++ api/rag/citation_res.json | 15 ++++ api/rag/main.py | 7 +- start_rag.sh | 33 +++++++ trustrag/modules/judger/chatgpt_judger.py | 38 +++----- trustrag/modules/rewriter/openai_rewrite.py | 97 ++++++++++++++------- 16 files changed, 272 insertions(+), 120 deletions(-) create mode 100644 api/README.md create mode 100644 api/rag/.env.example create mode 100644 api/rag/apps/config/base_config.py create mode 100644 api/rag/citation.json create mode 100644 api/rag/citation_res.json create mode 100644 start_rag.sh diff --git a/Dockerfile b/Dockerfile index 46970ee..ca339ae 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,5 +1,5 @@ # Use the official Ubuntu base image -FROM pytorch/pytorch:2.6.0-cuda12.6-cudnn9-runtime +FROM pytorch/pytorch:2.6.0-cuda12.6-cudnn9-devel ENV DEBIAN_FRONTEND=noninteractive ENV CUDA_DEVICE_ORDER=PCI_BUS_ID ENV PYTORCH_NVML_BASED_CUDA_CHECK=1 diff --git a/api/README.md b/api/README.md new file mode 100644 index 0000000..906dc55 --- /dev/null +++ b/api/README.md @@ -0,0 +1,24 @@ +# 使用默认 .env,再覆盖 RERANKER_NAME +sh start_rag.sh -e RERANKER_NAME=G:/pretrained_models/mteb/bge-reranker-large-v2 + + +docker run --rm \ + --env-file /path/to/.env \ + -e RERANKER_NAME=G:/pretrained_models/mteb/bge-reranker-large-v2 \ + -p 10000:10000 \ + -v /g/Projects/TrustRAG:/app \ + -w /app/api/rag \ + trustrag:v0.1 \ + sh -c "PYTHONPATH=/app python main.py" + + + +docker run --rm \ + --gpus all \ + -p 10000:10000 \ + -v /mnt/g/pretrained_models/mteb/bge-reranker-large:/mnt/g/pretrained_models/mteb/bge-reranker-large \ + -v .:/app \ + -e RERANKER_NAME=/mnt/g/pretrained_models/mteb/bge-reranker-large \ + -w /app/api/rag \ + trustrag:v0.1 \ + sh -c "PYTHONPATH=/app python main.py" diff --git a/api/rag/.env.example b/api/rag/.env.example new file mode 100644 index 0000000..7c78285 --- /dev/null +++ b/api/rag/.env.example @@ -0,0 +1,32 @@ +### API server +API_HOST=0.0.0.0 +API_PORT=10000 +API_URL=http://127.0.0.1:10001 + +### Upstream LLM service (gomall) +GOMALL_BASE_URL=http://10.208.61.1:32004/api/v1/1504_gomall_qwen3/Qwen3-30B-A3B-Instruct-2507/ +GOMALL_API_KEY=sk-xx + +### Local rerank service +RERANK_BASE_URL= +RERANK_API_KEY= + +### Model configs +LLM_NAME=Qwen3-30B-A3B-Instruct-2507 +LLM_PATH=G:/pretrained_models/llm/glm-4-9b-chat +LLM_SERVICE=online + +EMBEDDING_NAME=G:/pretrained_models/mteb/bge-large-zh-v1.5 +EMBEDDING_PATH=G:/pretrained_models/mteb/bge-large-zh-v1.5 +EMBEDDING_SERVICE=local + +RERANKER_NAME=G:/pretrained_models/mteb/bge-reranker-large +RERANKER_PATH=G:/pretrained_models/mteb/bge-reranker-large +RERANKER_SERVICE=local + +### Paths +DOCS_PATH=G:/Projects/TrustRAG/data/docs +INDEX_PATH=G:/Projects/TrustRAG/examples/retrievers/dense_cache + +### Rewriter tool +REWRITER_API_URL=http://10.208.61.1:32004/api/v1/1504_gomall_qwen3/Qwen3-30B-A3B-Instruct-2507/ \ No newline at end of file diff --git a/api/rag/apps/config/__init__.py b/api/rag/apps/config/__init__.py index fa2a8c0..4712ac4 100644 --- a/api/rag/apps/config/__init__.py +++ b/api/rag/apps/config/__init__.py @@ -1,11 +1,3 @@ -#!/usr/bin/env python -# -*- coding:utf-8 _*- -""" -@author:quincy qiang -@license: Apache Licence -@file: __init__.py -@time: 2024/06/13 -@contact: yanqiangmiffy@gamil.com -@software: PyCharm -@description: coding.. -""" +from .base_config import RAGSettings + +settings = RAGSettings() diff --git a/api/rag/apps/config/app_config.py b/api/rag/apps/config/app_config.py index c8c5a07..c12a52c 100644 --- a/api/rag/apps/config/app_config.py +++ b/api/rag/apps/config/app_config.py @@ -1,20 +1,9 @@ #!/usr/bin/env python # -*- coding:utf-8 _*- """ -@author:quincy qiang -@license: Apache Licence -@file: app_config.py.py -@time: 2024/06/13 -@contact: yanqiangmiffy@gamil.com -@software: PyCharm -@description: coding.. +应用基础配置。 """ -import pprint -from typing import ClassVar - -pp = pprint.PrettyPrinter(indent=4) - - +from api.rag.apps.config import settings class AppConfig: """配置类""" API_V1_STR: str = "" @@ -37,13 +26,13 @@ class AppConfig: {"url": "/v2", "description": "测试地址"}, ] - WEB_URL: ClassVar[str] = '*' + WEB_URL: str = '*' # 接口地址 - API_URL: ClassVar[str] = 'http://127.0.0.1:10001' + API_URL: str = settings.api_url # 运行访问的地址 - API_HOST: ClassVar[str] = '0.0.0.0' + API_HOST: str = settings.api_host # 端口 - API_PORT: int = 10000 + API_PORT: int = settings.api_port DEBUGGER: bool = True diff --git a/api/rag/apps/config/base_config.py b/api/rag/apps/config/base_config.py new file mode 100644 index 0000000..7ed0ee5 --- /dev/null +++ b/api/rag/apps/config/base_config.py @@ -0,0 +1,40 @@ +# from pydantic import BaseSettings +from pydantic import Field +from pydantic_settings import BaseSettings, SettingsConfigDict + + +class RAGSettings(BaseSettings): + """Centralized application settings loaded from environment variables.""" + + # API service + api_host: str = Field(default="0.0.0.0", validation_alias="API_HOST") + api_port: int = Field(default=10000, validation_alias="API_PORT") + api_url: str = Field(default="http://127.0.0.1:10001", validation_alias="API_URL") + + # Upstream LLM service (gomall) + gomall_base_url: str = Field(default="http://10.208.61.1:32004/api/v1/1504_gomall_qwen3/Qwen3-30B-A3B-Instruct-2507/", validation_alias="GOMALL_BASE_URL") + gomall_api_key: str = Field(default="", validation_alias="GOMALL_API_KEY") + llm_name: str = Field(default="Qwen3-30B-A3B-Instruct-2507", validation_alias="LLM_NAME") + # Tool specific + rewriter_api_url: str = Field(default="http://10.208.61.1:32004/api/v1/1504_gomall_qwen3/Qwen3-30B-A3B-Instruct-2507/", validation_alias="REWRITER_API_URL") + + # Local rerank service + rerank_base_url: str | None = Field(default=None, validation_alias="RERANK_BASE_URL") + rerank_api_key: str | None = Field(default=None, validation_alias="RERANK_API_KEY") + + + embedding_name: str = Field(default="G:/pretrained_models/mteb/bge-large-zh-v1.5", validation_alias="EMBEDDING_NAME") + embedding_path: str = Field(default="G:/pretrained_models/mteb/bge-large-zh-v1.5", validation_alias="EMBEDDING_PATH") + + reranker_name: str = Field(default="G:/pretrained_models/mteb/bge-reranker-large", validation_alias="RERANKER_NAME") + reranker_path: str = Field(default="G:/pretrained_models/mteb/bge-reranker-large", validation_alias="RERANKER_PATH") + # Paths + docs_path: str = Field(default="G:/Projects/TrustRAG/data/docs", validation_alias="DOCS_PATH") + index_path: str = Field(default="G:/Projects/TrustRAG/examples/retrievers/dense_cache", validation_alias="INDEX_PATH") + + model_config = SettingsConfigDict( + env_file=".env", + env_file_encoding="utf-8", + case_sensitive=False, + extra="ignore", + ) \ No newline at end of file diff --git a/api/rag/apps/config/rerank_config.py b/api/rag/apps/config/rerank_config.py index 7850011..b209f1c 100644 --- a/api/rag/apps/config/rerank_config.py +++ b/api/rag/apps/config/rerank_config.py @@ -9,17 +9,11 @@ @software: PyCharm @description: coding.. """ -from trustrag.config.config_loader import config +from api.rag.apps.config import settings -class RerankConfig(): +class RerankConfig: """重排序配置类""" - # 从配置文件加载服务和模型配置 - _rerank_service = config.get_config('services.rerank') - _rerank_model = config.get_config('models.reranker') - - # 模型名称 - model_name_or_path:str = _rerank_model['name'] - # 服务 URL - base_url:str = _rerank_service['base_url'] - # API 密钥 - api_key:str = _rerank_service['api_key'] + + model_name_or_path: str = settings.reranker_name + base_url: str | None = settings.rerank_base_url + api_key: str | None = settings.rerank_api_key diff --git a/api/rag/apps/core/judge/views.py b/api/rag/apps/core/judge/views.py index 37be325..2c87564 100644 --- a/api/rag/apps/core/judge/views.py +++ b/api/rag/apps/core/judge/views.py @@ -11,31 +11,26 @@ """ import loguru from fastapi import APIRouter -from trustrag.config.config_loader import ConfigLoader + +from api.rag.apps.config import settings from api.rag.apps.core.judge.bodys import JudgeBody from api.rag.apps.handle.response.json_response import ApiResponse from trustrag.modules.judger.bge_judger import BgeJudger, BgeJudgerConfig from trustrag.modules.judger.chatgpt_judger import OpenaiJudger, OpenaiJudgerConfig judge_router = APIRouter() -# 使用本地配置文件初始化配置加载器(单例) -config = ConfigLoader(config_path="config_local.json") -# 加载服务和模型配置 -llm_service = config.get_config('services.gomall') -llm_model = config.get_config('models.llm') -rerank_model = config.get_config('models.reranker') # BGE 判断器配置 judge_config = BgeJudgerConfig( - model_name_or_path=rerank_model['name'] + model_name_or_path=settings.reranker_name, ) bge_judger = BgeJudger(judge_config) # LLM 判断器配置 judger_config = OpenaiJudgerConfig( - base_url=llm_service['base_url'], - api_key=llm_service['api_key'], - model_name=llm_model['name'] + base_url=settings.gomall_base_url, + api_key=settings.gomall_api_key, + model_name=settings.llm_name, ) openai_judger = OpenaiJudger(judger_config) diff --git a/api/rag/apps/core/rerank/views.py b/api/rag/apps/core/rerank/views.py index e235b3b..380df04 100644 --- a/api/rag/apps/core/rerank/views.py +++ b/api/rag/apps/core/rerank/views.py @@ -12,22 +12,18 @@ import loguru from fastapi import APIRouter +from api.rag.apps.config import settings from api.rag.apps.core.rerank.bodys import RerankBody from api.rag.apps.core.rerank.models import Application from api.rag.apps.handle.response.json_response import UserNotFoundResponse, ApiResponse from trustrag.modules.reranker.bge_reranker import BgeReranker, BgeRerankerConfig -from trustrag.config.config_loader import ConfigLoader -config = ConfigLoader(config_path="config_local.json") rerank_router = APIRouter() -# 从配置文件加载重排序配置 -rerank_service = config.get_config('services.rerank') -rerank_model = config.get_config('models.reranker') reranker_config = BgeRerankerConfig( - model_name_or_path=rerank_model['name'], - api_key=rerank_service['api_key'], - url=rerank_service['base_url'] + model_name_or_path=settings.reranker_name, + api_key=settings.rerank_api_key, + url=settings.rerank_base_url, ) bge_reranker = BgeReranker(reranker_config) # Create diff --git a/api/rag/apps/core/rewrite/views.py b/api/rag/apps/core/rewrite/views.py index 501687f..50add4d 100644 --- a/api/rag/apps/core/rewrite/views.py +++ b/api/rag/apps/core/rewrite/views.py @@ -12,15 +12,18 @@ import loguru from fastapi import APIRouter +from api.rag.apps.config import settings from api.rag.apps.core.rewrite.bodys import RewriteBody from api.rag.apps.handle.response.json_response import ApiResponse -from trustrag.modules.rewriter.openai_rewrite import OpenaiRewriter,OpenaiRewriterConfig +from trustrag.modules.rewriter.openai_rewrite import OpenaiRewriter, OpenaiRewriterConfig rewriter_router = APIRouter() rewriter_config = OpenaiRewriterConfig( - api_url="http://10.208.63.29:8888" + base_url=settings.rewriter_api_url, + api_key=settings.gomall_api_key, + model_name=settings.llm_name, ) openai_rewriter = OpenaiRewriter(rewriter_config) diff --git a/api/rag/citation.json b/api/rag/citation.json new file mode 100644 index 0000000..fd64036 --- /dev/null +++ b/api/rag/citation.json @@ -0,0 +1,13 @@ +{ + "question": "请介绍下巨齿鲨2电影", + "response": "巨齿鲨2是一部科幻冒险电影,由本·维特利执导,杰森·斯坦森、吴京、蔡书雅和克利夫·柯蒂斯主演。电影讲述了海洋霸主巨齿鲨,今夏再掀狂澜!乔纳斯·泰勒(杰森·斯坦森饰)与科学家张九溟(吴京饰)双雄联手,进入海底7000米深渊执行探索任务。他们意外遭遇史前巨兽海洋霸主巨齿鲨群的攻击,还将对战凶猛危险的远古怪兽群。惊心动魄的深渊冒险,巨燃巨爽的深海大战一触即发。", + "evidences": [ + "海洋霸主巨齿鲨,今夏再掀狂澜!乔纳斯·泰勒(杰森·斯坦森 饰)与科学家张九溟(吴京 饰)双雄联手,进入海底7000米深渊执行探索任务。他们意外遭遇史前巨兽海洋霸主巨齿鲨群的攻击,还将对战凶猛危险的远古怪兽群。惊心动魄的深渊冒险,巨燃巨爽的深海大战一触即发", + "本·维特利 编剧:乔·霍贝尔埃里希·霍贝尔迪恩·乔格瑞斯 国家地区:中国 | 美国 发行公司:上海华人影业有限公司五洲电影发行有限公司中国电影股份有限公司北京电影发行分公司 出品公司:上海华人影业有限公司华纳兄弟影片公司北京登峰国际文化传播有限公司 更多片名:巨齿鲨2 剧情:海洋霸主巨齿鲨,今夏再掀狂澜!乔纳斯·泰勒(杰森·斯坦森 饰)与科学家张九溟(吴京 饰)双雄联手,进入海底7000米深渊执行探索任务。他们意外遭遇史前巨兽海洋霸主巨齿鲨群的攻击,还将对战凶猛危险的远古怪兽群。惊心动魄的深渊冒险,巨燃巨爽的深海大战一触即发……" + ], + "selected_idx": [ + 1, + 2 + ], + "selected_docs": [] +} \ No newline at end of file diff --git a/api/rag/citation_res.json b/api/rag/citation_res.json new file mode 100644 index 0000000..c99d8f1 --- /dev/null +++ b/api/rag/citation_res.json @@ -0,0 +1,15 @@ +{ + "question": "请介绍下巨齿鲨2电影", + "response": "巨齿鲨2是一部科幻冒险电影,由本·维特利执导,杰森·斯坦森、吴京、蔡书雅和克利夫·柯蒂斯主演。电影讲述了海洋霸主巨齿鲨,今夏再掀狂澜!乔纳斯·泰勒(杰森·斯坦森饰)与科学家张九溟(吴京饰)双雄联手,进入海底7000米深渊执行探索任务。他们意外遭遇史前巨兽海洋霸主巨齿鲨群的攻击,还将对战凶猛危险的远古怪兽群。惊心动魄的深渊冒险,巨燃巨爽的深海大战一触即发。", + "evidences": [ + "海洋霸主巨齿鲨,今夏再掀狂澜!乔纳斯·泰勒(杰森·斯坦森 饰)与科学家张九溟(吴京 饰)双雄联手,进入海底7000米深渊执行探索任务。他们意外遭遇史前巨兽海洋霸主巨齿鲨群的攻击,还将对战凶猛危险的远古怪兽群。惊心动魄的深渊冒险,巨燃巨爽的深海大战一触即发", + "本·维特利 编剧:乔·霍贝尔埃里希·霍贝尔迪恩·乔格瑞斯 国家地区:中国 | 美国 发行公司:上海华人影业有限公司五洲电影发行有限公司中国电影股份有限公司北京电影发行分公司 出品公司:上海华人影业有限公司华纳兄弟影片公司北京登峰国际文化传播有限公司 更多片名:巨齿鲨2 剧情:海洋霸主巨齿鲨,今夏再掀狂澜!乔纳斯·泰勒(杰森·斯坦森 饰)与科学家张九溟(吴京 饰)双雄联手,进入海底7000米深渊执行探索任务。他们意外遭遇史前巨兽海洋霸主巨齿鲨群的攻击,还将对战凶猛危险的远古怪兽群。惊心动魄的深渊冒险,巨燃巨爽的深海大战一触即发……" + ], + "selected_idx": [ + 1, + 2 + ], + "selected_docs": [], + "result": "", + "quote_list": [] +} \ No newline at end of file diff --git a/api/rag/main.py b/api/rag/main.py index e3c564c..3ead246 100644 --- a/api/rag/main.py +++ b/api/rag/main.py @@ -11,11 +11,10 @@ """ import os import sys - +from dotenv import load_dotenv +load_dotenv() os.environ["TOKENIZERS_PARALLELISM"] = "false" -sys.path.append('.') -# sys.path.append('/data/users/searchgpt/yq/GoMate') -sys.path.append('/data/users/searchgpt/yq/GoMate_dev') +sys.path.append('../../') sys.path.append('/home/yanqiang/code') from apps.app import create_app diff --git a/start_rag.sh b/start_rag.sh new file mode 100644 index 0000000..883bf0a --- /dev/null +++ b/start_rag.sh @@ -0,0 +1,33 @@ +#!/usr/bin/env sh +set -eu + +# Start the RAG API (api/rag/main.py) inside the trustrag:v0.1 image. +# POSIX 兼容,sh 或 bash 均可执行。 +# 用法: +# sh start_rag.sh +# ENV_FILE=path/to/.env HOST_PORT=18000 IMAGE_NAME=trustrag:v0.1 sh start_rag.sh + +ROOT_DIR=$(cd "$(dirname "$0")" && pwd) +ENV_FILE=${ENV_FILE:-"$ROOT_DIR/api/rag/.env"} +HOST_PORT=${HOST_PORT:-10000} +CONTAINER_PORT=10000 +IMAGE_NAME=${IMAGE_NAME:-trustrag:v0.1} + +if [ ! -f "$ENV_FILE" ]; then + echo "[WARN] Env file not found at $ENV_FILE, continuing without --env-file" + ENV_ARG="" +else + ENV_ARG="--env-file $ENV_FILE" +fi + +# 挂载整个仓库,确保能找到 trustrag 包 +docker run --rm \ + -p "${HOST_PORT}:${CONTAINER_PORT}" \ + $ENV_ARG \ + -v "$ROOT_DIR":/app \ + -w /app/api/rag \ + "$IMAGE_NAME" \ + sh -c "PYTHONPATH=/app python main.py" + + + diff --git a/trustrag/modules/judger/chatgpt_judger.py b/trustrag/modules/judger/chatgpt_judger.py index 70ddbb1..fad4b08 100644 --- a/trustrag/modules/judger/chatgpt_judger.py +++ b/trustrag/modules/judger/chatgpt_judger.py @@ -1,7 +1,7 @@ import time from typing import List, Any -import requests +from openai import OpenAI from tqdm import tqdm from trustrag.modules.judger.base import BaseJudger @@ -37,8 +37,11 @@ class OpenaiJudger(BaseJudger): def __init__(self, config): super().__init__() self.config = config - self.base_url = config.base_url - self.api_key = config.api_key + self.client = OpenAI( + base_url=config.base_url, + api_key=config.api_key or "", + timeout=30, + ) self.model_name = config.model_name print('成功初始化 ChatGPT 判断器') @@ -58,31 +61,18 @@ def judge(self, query: str, documents: List[str], k: int = 5, is_sorted: bool = 注意:只返回 1 或 0,不解释原因,不输出其他内容。 """ - headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {self.api_key}" - } - results = [] for doc in tqdm(documents, desc="判断文档相关性"): - data = { - "model": self.model_name, - "messages": [ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": f"查询:{query}\n\n文章:{doc}"} - ] - } - try: - response = requests.post( - self.base_url + "/chat/completions", - headers=headers, - json=data, - timeout=30 + completion = self.client.chat.completions.create( + model=self.model_name, + messages=[ + {"role": "user", "content": system_prompt+f"查询:{query}\n\n文章:{doc}"}, + ], + temperature=0.0, ) - response.raise_for_status() - result = response.json() - score = float(result['choices'][0]['message']['content'].strip()) + result_text = completion.choices[0].message.content.strip() + score = float(result_text) results.append({"text": doc, "score": score}) except Exception as e: print(f"调用 LLM 服务失败: {str(e)}") diff --git a/trustrag/modules/rewriter/openai_rewrite.py b/trustrag/modules/rewriter/openai_rewrite.py index ed0f2c0..f8f1953 100644 --- a/trustrag/modules/rewriter/openai_rewrite.py +++ b/trustrag/modules/rewriter/openai_rewrite.py @@ -1,39 +1,76 @@ -import time -from typing import List, Any -import requests -from tqdm import tqdm -from trustrag.modules.rewriter.base import BaseRewriter import json import re +from typing import List, Any + +from openai import OpenAI +from trustrag.modules.rewriter.base import BaseRewriter class OpenaiRewriterConfig: - """ - """ + """Config for OpenAI-compatible rewriter.""" - def __init__(self, api_url='http://gomatellm-service.aicloud-yanqiang.svc.cluster.local'): - self.api_url = api_url + def __init__(self, base_url: str, api_key: str | None = None, model_name: str | None = None, timeout: int = 30): + self.base_url = base_url + self.api_key = api_key + self.model_name = model_name + self.timeout = timeout def log_config(self): - # Log the current configuration settings return f""" - BgeRewriterConfig: - API URL: {self.api_url} + OpenaiRewriterConfig: + Base URL: {self.base_url} + API Key: {'*' * 8 if self.api_key else 'Not Set'} + Model Name: {self.model_name} + Timeout: {self.timeout}s """ class OpenaiRewriter(BaseRewriter): """ - A Rewriter that utilizes a BERT-based model for sequence classification - to judge a list of documents based on their relevance to a given query. + A Rewriter that calls an OpenAI-compatible chat completion endpoint. """ - def __init__(self, config): + def __init__(self, config: OpenaiRewriterConfig): super().__init__() self.config = config - self.api_url = self.config.api_url + self.client = OpenAI( + base_url=self.config.base_url, + api_key=self.config.api_key or "", + timeout=self.config.timeout, + ) + self.model_name = self.config.model_name print('Successful Init ChatGPT Rewriter ') + def repair_json_output(self,content: str) -> str: + """ + Repair and normalize JSON output. + + Args: + content (str): String content that may contain JSON + + Returns: + str: Repaired JSON string, or original content if not JSON + """ + content = content.strip() + if content.startswith(("{", "[")) or "```json" in content or "```ts" in content: + try: + # If content is wrapped in ```json code block, extract the JSON part + if content.startswith("```json"): + content = content.removeprefix("```json") + + if content.startswith("```ts"): + content = content.removeprefix("```ts") + + if content.endswith("```"): + content = content.removesuffix("```") + + # Try to repair and parse JSON + repaired_content = json.loads(content) + return json.dumps(repaired_content, ensure_ascii=False) + except Exception as e: + print(f"JSON repair failed: {e}") + return content + def parse_response(self, response_data: str): """ 解析JSON响应字符串,如果解析失败则返回默认空值 @@ -54,10 +91,11 @@ def parse_response(self, response_data: str): # 如果response是字符串,则需要再次解析 if isinstance(response_data, str): try: - response_data = re.sub(r'^.*?```json\n|```$', '', response_data, flags=re.DOTALL) + # response_data = re.sub(r'^.*?```json\n|```$', '', response_data, flags=re.DOTALL) + response_data=self.repair_json_output(response_data) response_data = json.loads(response_data) except json.JSONDecodeError: - print("报错") + print("报错",response_data) return default_response # 从解析后的数据中提取字段,如果不存在则使用空字符串 @@ -76,7 +114,7 @@ def parse_response(self, response_data: str): print("报错") return default_response - def rewrite(self, query): + def rewrite(self, query: str) -> dict: system_prompt = """ 请分析用户问题并提取其中的地点、时间、活动或会议名称,将这些信息以JSON格式输出。如果信息不全或用户未提及,则标记为""。按以下格式生成JSON输出: { @@ -130,15 +168,14 @@ def rewrite(self, query): 用户问题: """ - # Request payload - payload = { - "prompt": system_prompt + "\n" + query, - "teampture": 0.2, - "top_k": 20 - } - print(system_prompt + "\n" + query) - response = requests.post(self.api_url, json=payload) - response = response.json() - response = self.parse_response(response['response']) - return response + completion = self.client.chat.completions.create( + model=self.model_name, + messages=[ + # {"role": "system", "content": system_prompt}, + {"role": "user", "content": system_prompt+query}, + ], + temperature=0.2, + ) + content = completion.choices[0].message.content + return self.parse_response(content)