init
This commit is contained in:
@@ -15,3 +15,7 @@ def get_service_config(request: Request):
|
||||
|
||||
def get_tool_router(request: Request):
|
||||
return request.app.state.tool_router
|
||||
|
||||
|
||||
def get_prompt_manager(request: Request):
|
||||
return request.app.state.prompt_manager
|
||||
|
||||
+45
-4
@@ -6,7 +6,8 @@ from schemas.agent_output import AgentOutput
|
||||
from schemas.tool_input import ToolInput
|
||||
from schemas.tool_output import ToolOutput
|
||||
from workflows.workflow_manager import WorkflowType
|
||||
from api.dependencies import get_workflow_manager, get_nacos_manager, get_service_config, get_tool_router
|
||||
from api.dependencies import get_workflow_manager, get_nacos_manager, get_service_config, get_tool_router, get_prompt_manager
|
||||
from services.ragflow_sync import RagflowSync
|
||||
|
||||
|
||||
router = APIRouter()
|
||||
@@ -63,12 +64,32 @@ def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workfl
|
||||
if workflow_type != WorkflowType.CONVERSATION:
|
||||
raise HTTPException(status_code=400, detail="仅支持对话工作流的流式输出")
|
||||
|
||||
agent = workflow_manager.get_workflow(workflow_type)
|
||||
def _extract_output_text(result: dict) -> str:
|
||||
context = (result.get("context") or {}) if isinstance(result, dict) else {}
|
||||
if "sr_api_result" in context:
|
||||
return str(context.get("sr_api_result") or "")
|
||||
messages = result.get("messages") if isinstance(result, dict) else None
|
||||
if messages:
|
||||
last = messages[-1]
|
||||
if hasattr(last, "content"):
|
||||
return str(last.content or "")
|
||||
return ""
|
||||
|
||||
def event_stream():
|
||||
try:
|
||||
for token in agent.stream_run(payload.input):
|
||||
yield f"data: {token}\n\n"
|
||||
result = workflow_manager.execute_workflow(
|
||||
workflow_type=workflow_type,
|
||||
user_input=payload.input,
|
||||
session_id=payload.session_id,
|
||||
)
|
||||
text = _extract_output_text(result.get("result") or {})
|
||||
if not text:
|
||||
yield "event: end\ndata: [DONE]\n\n"
|
||||
return
|
||||
chunk_size = 512
|
||||
for i in range(0, len(text), chunk_size):
|
||||
chunk = text[i : i + chunk_size]
|
||||
yield f"data: {chunk}\n\n"
|
||||
yield "event: end\ndata: [DONE]\n\n"
|
||||
except Exception as e:
|
||||
yield f"event: error\ndata: {str(e)}\n\n"
|
||||
@@ -80,3 +101,23 @@ def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workfl
|
||||
def run_tool(payload: ToolInput, tool_router=Depends(get_tool_router)):
|
||||
result = tool_router.call(payload.tool_name, payload.payload)
|
||||
return ToolOutput(**result)
|
||||
|
||||
|
||||
@router.post("/api/prompts/reload")
|
||||
def reload_prompts(prompt_manager=Depends(get_prompt_manager)):
|
||||
prompt_manager.reload()
|
||||
return {"ok": True}
|
||||
|
||||
|
||||
@router.post("/api/ragflow/table-retrieval/reload")
|
||||
def reload_table_retrieval():
|
||||
syncer = RagflowSync()
|
||||
result = syncer.sync_table_retrieval()
|
||||
return {"ok": True, "result": result}
|
||||
|
||||
|
||||
@router.post("/api/ragflow/sql-gen/reload")
|
||||
def reload_sql_gen():
|
||||
syncer = RagflowSync()
|
||||
result = syncer.sync_sql_gen_prompts()
|
||||
return {"ok": True, "result": result}
|
||||
|
||||
Reference in New Issue
Block a user