Commit c7bc7aac by ccran

feat: rename retrieve reference; add nanobot llm;

parent 584b5dc7
......@@ -118,6 +118,9 @@ LLM = {
base_url=f"{base_fastgpt_url}/api/v1", api_key=reflect_retry_api_key
)
),
"nanobot_llm": LLMConfig(
base_url="http://172.21.107.80:19090/v1", api_key='none', model="Qwen3.5-122B-A10B-AWQ"
),
}
doc_support_formats = [".docx", ".doc", ".wps"]
pdf_support_formats = [".txt", ".md", ".pdf"]
......@@ -9,8 +9,8 @@ from core.tool import ToolBase, tool, tool_func
from utils.excel_util import ExcelUtil
@tool("retrieve_reference", "审查参考检索")
class RetrieveReferenceTool(ToolBase):
@tool("rules_retrieve_reference", "审查参考检索")
class RulesRetrieveReferenceTool(ToolBase):
def __init__(self) -> None:
self.default_ruleset_id = DEFAULT_RULESET_ID
self.column_map = {
......@@ -72,7 +72,7 @@ class RetrieveReferenceTool(ToolBase):
if __name__ == "__main__":
tool = RetrieveReferenceTool()
tool = RulesRetrieveReferenceTool()
result = tool.run(ruleset_id="金盘", routed_rule_titles=None)
for rule in result.get("rules", []):
print(f"Rule Title: {rule.get('title')}")
......
......@@ -13,17 +13,20 @@ class LLMTool(ToolBase):
"""LLM-backed processor: builds prompts, calls LLM, parses JSON."""
def __init__(
self, system_prompt: str, llm_key: str = "fastgpt_segment_review"
self, system_prompt: str = None, llm_key: str = "fastgpt_segment_review"
) -> None:
super().__init__()
self.system_prompt = system_prompt
self.llm = OpenAITool(LLM[llm_key], max_workers=MAX_WORKERS)
def build_messages(self, user_content: str, system_content: str = None) -> List[Dict[str, str]]:
if system_content or self.system_prompt:
return [
{"role": "system", "content": system_content or self.system_prompt},
{"role": "user", "content": user_content},
]
else:
return [{"role": "user", "content": user_content}]
async def chat_async(self, messages: List[Dict[str, str]]):
return await self.llm.chat(messages)
......@@ -47,3 +50,8 @@ class LLMTool(ToolBase):
return data[0] if data else {}
except Exception:
return {}
if __name__ == "__main__":
tool = LLMTool(llm_key="nanobot_llm")
results = asyncio.run(tool.chat_async(tool.build_messages("列出工作目录")))
print(results)
\ No newline at end of file
......@@ -26,7 +26,7 @@ from core.tools.segment_summary import SegmentSummaryTool
from core.tools.segment_review import SegmentReviewTool
from core.tools.segment_rule_router import SegmentRuleRouterTool
from core.tools.rule_filter import LufaPartyRuleFilterTool
from core.tools.retrieve_reference import RetrieveReferenceTool
from core.tools.rules_retrieve_reference import RulesRetrieveReferenceTool
from core.tools.reflect_retry import ReflectRetryTool
from core.tools.segment_merger import SegmentMergerTool
from core.tools.fact_merger import FactMergerTool
......@@ -42,7 +42,7 @@ summary_tool = SegmentSummaryTool()
review_tool = SegmentReviewTool()
rule_router_tool = SegmentRuleRouterTool()
lufa_party_rule_filter_tool = LufaPartyRuleFilterTool()
reference_tool = RetrieveReferenceTool()
rules_reference_tool = RulesRetrieveReferenceTool()
reflect_tool = ReflectRetryTool()
merger_tool = SegmentMergerTool()
fact_merger_tool = FactMergerTool()
......@@ -74,6 +74,7 @@ class DocumentParseResponse(BaseModel):
ruleset_items: List[str]
summary_names: List[str]
text: Optional[str] = None
text_file_path: Optional[str] = None
file_ext: Optional[str] = None
file_name: Optional[str] = None
......@@ -155,12 +156,14 @@ async def parse_document(payload: DocumentParseRequest) -> DocumentParseResponse
# ocr
await doc_obj.get_from_ocr()
text = doc_obj.get_all_text()
text_file_path = Path(file_path).with_suffix(".txt").resolve()
text_file_path.write_text(text or "", encoding="utf-8")
segment_ids = doc_obj.get_chunk_id_list()
# TODO: FastGPT BUG segment_ids必须从1开始,0开始会缺少第一段文本,后续需要修复
segment_ids = [idx + 1 for idx in segment_ids]
# get ruleset items
ruleset_id = payload.ruleset_id or reference_tool.default_ruleset_id
ruleset_items = reference_tool.run(ruleset_id=ruleset_id).get("rules", [])
ruleset_id = payload.ruleset_id or rules_reference_tool.default_ruleset_id
ruleset_items = rules_reference_tool.run(ruleset_id=ruleset_id).get("rules", [])
ruleset_review_items = [
t
for t in (r.get("title") for r in ruleset_items)
......@@ -169,13 +172,14 @@ async def parse_document(payload: DocumentParseRequest) -> DocumentParseResponse
summary_names = list(
dict.fromkeys(
s.strip()
for s in (r.get("summary") for r in ruleset_items)
for s in (r.get("summa y") for r in ruleset_items)
if isinstance(s, str) and s.strip()
)
)
return DocumentParseResponse(
conversation_id=payload.conversation_id,
text=text,
text_file_path=str(text_file_path),
segment_ids=segment_ids,
ruleset_items=ruleset_review_items,
summary_names=summary_names,
......@@ -223,21 +227,21 @@ def summarize_facts(payload: SegmentSummaryRequest) -> SegmentSummaryResponse:
detail=f"Segment text not found for id {payload.segment_id}: {exc}. Please parse document first.",
)
ruleset_id = payload.ruleset_id or reference_tool.default_ruleset_id
ruleset_id = payload.ruleset_id or rules_reference_tool.default_ruleset_id
if payload.routed_summary_names is not None:
summary_names = {
name.strip()
for name in payload.routed_summary_names
if isinstance(name, str) and name.strip()
}
all_rules = reference_tool.run(ruleset_id=ruleset_id).get("rules", [])
all_rules = rules_reference_tool.run(ruleset_id=ruleset_id).get("rules", [])
rules = [
rule
for rule in all_rules
if str(rule.get("summary", "")).strip() in summary_names
]
else:
rules = reference_tool.run(
rules = rules_reference_tool.run(
ruleset_id=ruleset_id,
routed_rule_titles=payload.routed_rule_titles,
).get("rules", [])
......@@ -309,13 +313,13 @@ def review_segment(payload: SegmentReviewRequest) -> SegmentReviewResponse:
detail=f"Segment text not found for id {payload.segment_id}: {exc}. Please parse document first.",
)
ruleset_id = payload.ruleset_id or reference_tool.default_ruleset_id
rules = reference_tool.run(
ruleset_id = payload.ruleset_id or rules_reference_tool.default_ruleset_id
rules = rules_reference_tool.run(
ruleset_id=ruleset_id,
routed_rule_titles=payload.routed_rule_titles,
).get("rules", [])
# 暂时不添加摘要看下结果
# summary_keywords = reference_tool.summary_keywords(rules)
# summary_keywords = rules_reference_tool.summary_keywords(rules)
# context_summaries = store.search_facts(summary_keywords)
result = review_tool.run(
......@@ -378,8 +382,8 @@ def route_segment_rules(payload: SegmentReviewRequest) -> SegmentRuleRouterRespo
detail=f"Segment text not found for id {payload.segment_id}: {exc}. Please parse document first.",
)
ruleset_id = payload.ruleset_id or reference_tool.default_ruleset_id
rules = reference_tool.run(ruleset_id=ruleset_id).get("rules", [])
ruleset_id = payload.ruleset_id or rules_reference_tool.default_ruleset_id
rules = rules_reference_tool.run(ruleset_id=ruleset_id).get("rules", [])
if use_lufa and rules:
try:
......@@ -458,8 +462,8 @@ class FactsMergerResponse(BaseModel):
@app.post("/segments/review/reflect", response_model=ReflectReviewResponse)
def reflect_review(payload: ReflectReviewRequest) -> ReflectReviewResponse:
store = get_cached_memory(payload.conversation_id)
ruleset_id = payload.ruleset_id or reference_tool.default_ruleset_id
ruleset_items = reference_tool.run(ruleset_id=ruleset_id).get("rules", [])
ruleset_id = payload.ruleset_id or rules_reference_tool.default_ruleset_id
ruleset_items = rules_reference_tool.run(ruleset_id=ruleset_id).get("rules", [])
rule = next(
(r for r in ruleset_items if r.get("title") == payload.rule_title), None
)
......@@ -467,7 +471,7 @@ def reflect_review(payload: ReflectReviewRequest) -> ReflectReviewResponse:
raise HTTPException(
status_code=404, detail=f"Rule not found: {payload.rule_title}"
)
summary_keywords = reference_tool.summary_keywords([rule])
summary_keywords = rules_reference_tool.summary_keywords([rule])
context_summaries_facts = store.search_facts(summary_keywords)
# 查找审查规则对应的 findings
findings = [
......
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or sign in to comment