From 711e3c55ca61025239c720b7db2b0d63b71ad25f Mon Sep 17 00:00:00 2001 From: chengke <404835780@qq.com> Date: Wed, 13 May 2026 12:14:52 +0800 Subject: [PATCH 1/8] feat: implement decomposed retrieval workflow system with query planning, synthesis, and execution budget management. --- apps/api/.env.example | 8 + apps/api/app/api/v1/routes/retrieval.py | 36 +- apps/worker/.env.example | 8 + .../shared-python/shared/core/config/ai.py | 32 ++ .../shared/models/database/document.py | 3 + .../shared/services/retrieval/__init__.py | 3 +- .../retrieval/agentic/orchestrator.py | 127 ++++-- .../services/retrieval/agentic/tools.py | 13 +- .../services/retrieval/agentic/trace.py | 13 + .../shared/services/retrieval/app_service.py | 62 +-- .../services/retrieval/cache_service.py | 37 ++ .../shared/services/retrieval/llm_adapter.py | 46 +++ .../services/retrieval/workflow/__init__.py | 13 + .../retrieval/workflow/orchestrator.py | 382 ++++++++++++++++++ .../services/retrieval/workflow/planner.py | 239 +++++++++++ .../retrieval/workflow/synthesizer.py | 125 ++++++ .../services/retrieval/workflow/types.py | 230 +++++++++++ .../services/retrieval/workflow/wallet.py | 150 +++++++ .../tests/test_workflow_orchestrator.py | 67 +++ .../shared/tests/test_workflow_planner.py | 89 ++++ .../shared/tests/test_workflow_wallet.py | 80 ++++ 21 files changed, 1684 insertions(+), 79 deletions(-) create mode 100644 packages/shared-python/shared/services/retrieval/workflow/__init__.py create mode 100644 packages/shared-python/shared/services/retrieval/workflow/orchestrator.py create mode 100644 packages/shared-python/shared/services/retrieval/workflow/planner.py create mode 100644 packages/shared-python/shared/services/retrieval/workflow/synthesizer.py create mode 100644 packages/shared-python/shared/services/retrieval/workflow/types.py create mode 100644 packages/shared-python/shared/services/retrieval/workflow/wallet.py create mode 100644 packages/shared-python/shared/tests/test_workflow_orchestrator.py create mode 100644 packages/shared-python/shared/tests/test_workflow_planner.py create mode 100644 packages/shared-python/shared/tests/test_workflow_wallet.py diff --git a/apps/api/.env.example b/apps/api/.env.example index 54c9a9709..4a01c6d8d 100644 --- a/apps/api/.env.example +++ b/apps/api/.env.example @@ -85,6 +85,14 @@ NORMOL_MODEL=deepseek-chat HIERARCHY_LLM_MODEL=qwen3.6-flash IMAGE_MODEL=qwen3.5-flash IMAGE_MODEL_MAX=qwen3.5-flash +RETRIEVAL_DECOMPOSITION_ENABLED=false +RETRIEVAL_PLANNER_MODEL= +RETRIEVAL_PLANNER_THINKING_BUDGET=4000 +RETRIEVAL_DECOMPOSITION_MAX_STEPS=5 +RETRIEVAL_WALLET_TOTAL_BUDGET=200000 +RETRIEVAL_WALLET_PER_RETRIEVE_STEP_BUDGET=40000 +RETRIEVAL_WALLET_PER_SYNTHESIZE_STEP_BUDGET=6000 +RETRIEVAL_WORKFLOW_PARALLEL_MAX=3 # File handling defaults SUPPORTED_EXTENSIONS=.doc,.docx,.pdf,.txt,.xls,.xlsx,.csv,.pptx,.jpg,.jpeg,.png,.md diff --git a/apps/api/app/api/v1/routes/retrieval.py b/apps/api/app/api/v1/routes/retrieval.py index e0106f0bc..aaee0b7bc 100644 --- a/apps/api/app/api/v1/routes/retrieval.py +++ b/apps/api/app/api/v1/routes/retrieval.py @@ -52,6 +52,11 @@ class RetrievalQueryRequest(BaseModel): internal_recall_k: int | None = Field( None, ge=1, description="Override per-channel recall count" ) + enable_decomposition: bool | None = Field( + None, + description="Deprecated: agentic mode now always uses workflow decomposition. This field is ignored.", + deprecated=True, + ) @field_validator("channels") @classmethod @@ -63,7 +68,35 @@ def validate_channels(cls, v: list[str]) -> list[str]: return v -@router.post("/query") +class WorkflowStepResponse(BaseModel): + step_id: str + sub_query: str + step_kind: Literal["retrieve", "synthesize"] + depends_on: list[str] + output_role: str + status: Literal["done", "skipped", "error", "budget_stop"] + answer_text: str + evidence_text: str | None = None + referenced_chunks: list[dict] = Field(default_factory=list) + budget_snapshot: dict | None = None + child_run_id: str | None = None + + +class RetrievalQueryResponse(BaseModel): + namespace: str + query: str + router_used: str + answer_text: str | None = None + referenced_chunks: list[dict] = Field(default_factory=list) + results: list[dict] = Field(default_factory=list) + plan: dict | None = None + steps: list[WorkflowStepResponse] | None = None + final_strategy_used: str | None = None + wallet_snapshot: dict | None = None + planner_snapshot: dict | None = None + + +@router.post("/query", response_model=RetrievalQueryResponse) async def query_retrieval( payload: RetrievalQueryRequest, current_user: CurrentUser = Depends(with_current_user), @@ -85,4 +118,5 @@ async def query_retrieval( rerank=payload.rerank, threshold=payload.threshold, internal_recall_k=payload.internal_recall_k, + enable_decomposition=payload.enable_decomposition, ) diff --git a/apps/worker/.env.example b/apps/worker/.env.example index 49bb7aae4..6c46d9df7 100644 --- a/apps/worker/.env.example +++ b/apps/worker/.env.example @@ -85,6 +85,14 @@ NORMOL_MODEL=deepseek-chat HIERARCHY_LLM_MODEL=deepseek-chat IMAGE_MODEL=qwen3.5-flash IMAGE_MODEL_MAX=qwen3.5-flash +RETRIEVAL_DECOMPOSITION_ENABLED=false +RETRIEVAL_PLANNER_MODEL= +RETRIEVAL_PLANNER_THINKING_BUDGET=4000 +RETRIEVAL_DECOMPOSITION_MAX_STEPS=5 +RETRIEVAL_WALLET_TOTAL_BUDGET=200000 +RETRIEVAL_WALLET_PER_RETRIEVE_STEP_BUDGET=40000 +RETRIEVAL_WALLET_PER_SYNTHESIZE_STEP_BUDGET=6000 +RETRIEVAL_WORKFLOW_PARALLEL_MAX=3 # Required for specific features: billing and analytics BILLING_ENABLED=false diff --git a/packages/shared-python/shared/core/config/ai.py b/packages/shared-python/shared/core/config/ai.py index 2b25b5571..2164df3d2 100644 --- a/packages/shared-python/shared/core/config/ai.py +++ b/packages/shared-python/shared/core/config/ai.py @@ -35,6 +35,38 @@ class AIConfig(BaseModel): default="qwen3.5-flash", description="Higher-capability image model for OCR and ask-image Q&A", ) + RETRIEVAL_DECOMPOSITION_ENABLED: bool = Field( + default=False, + description="Enable query-decomposition workflow before agentic retrieval.", + ) + RETRIEVAL_PLANNER_MODEL: str = Field( + default="", + description="Reasoning-capable model used by the workflow query planner.", + ) + RETRIEVAL_PLANNER_THINKING_BUDGET: int = Field( + default=4000, + description="Token budget for the query planner thinking call.", + ) + RETRIEVAL_DECOMPOSITION_MAX_STEPS: int = Field( + default=5, + description="Maximum number of planned workflow steps.", + ) + RETRIEVAL_WALLET_TOTAL_BUDGET: int = Field( + default=200000, + description="Total workflow token wallet for decomposed retrieval.", + ) + RETRIEVAL_WALLET_PER_RETRIEVE_STEP_BUDGET: int = Field( + default=40000, + description="Default token budget issued to each retrieve step.", + ) + RETRIEVAL_WALLET_PER_SYNTHESIZE_STEP_BUDGET: int = Field( + default=6000, + description="Default token budget issued to each synthesize step.", + ) + RETRIEVAL_WORKFLOW_PARALLEL_MAX: int = Field( + default=3, + description="Maximum concurrent workflow steps in the same DAG batch.", + ) # Runtime LLM controls. LLM_MOCK_ENABLED: bool = Field( diff --git a/packages/shared-python/shared/models/database/document.py b/packages/shared-python/shared/models/database/document.py index 8b5ba0a23..41bed45bb 100644 --- a/packages/shared-python/shared/models/database/document.py +++ b/packages/shared-python/shared/models/database/document.py @@ -380,6 +380,9 @@ class RetrievalRun(Base): result_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) final_doc_ids: Mapped[Optional[List[str]]] = mapped_column(JSON, nullable=True) result_provenance: Mapped[Optional[Dict[str, Any]]] = mapped_column(JSON, nullable=True) + parent_run_id: Mapped[Optional[str]] = mapped_column(String(36), nullable=True, index=True) + workflow_step_id: Mapped[Optional[str]] = mapped_column(String(64), nullable=True) + workflow_plan: Mapped[Optional[Dict[str, Any]]] = mapped_column(JSON, nullable=True) latency_ms: Mapped[int] = mapped_column(Integer, nullable=False, default=0) token_count: Mapped[Optional[int]] = mapped_column(Integer, nullable=True) error: Mapped[Optional[str]] = mapped_column(Text, nullable=True) diff --git a/packages/shared-python/shared/services/retrieval/__init__.py b/packages/shared-python/shared/services/retrieval/__init__.py index a25d45cdc..832e81eb6 100644 --- a/packages/shared-python/shared/services/retrieval/__init__.py +++ b/packages/shared-python/shared/services/retrieval/__init__.py @@ -8,10 +8,11 @@ ) from .graph_service import DocumentGraphService, GraphQueryService, GraphScope from .hit_stats_service import record_retrieval_hits -from .llm_adapter import create_retrieval_llm_fn +from .llm_adapter import create_retrieval_llm_fn, create_retrieval_planner_fn __all__ = [ "create_retrieval_llm_fn", + "create_retrieval_planner_fn", "run_retrieval_query", "merge_channels_rrf", "DocumentGraphService", diff --git a/packages/shared-python/shared/services/retrieval/agentic/orchestrator.py b/packages/shared-python/shared/services/retrieval/agentic/orchestrator.py index b8b766327..3bc3b7774 100644 --- a/packages/shared-python/shared/services/retrieval/agentic/orchestrator.py +++ b/packages/shared-python/shared/services/retrieval/agentic/orchestrator.py @@ -100,6 +100,65 @@ async def _build_asset_url_map( return url_map +def _collect_all_leaf_paths(node: DocTreeNode) -> set[str]: + """Recursively collect all leaf_content keys across the entire tree.""" + paths = set(node.leaf_content.keys()) + for child in node.children.values(): + paths.update(_collect_all_leaf_paths(child)) + return paths + + +def _reconcile_deferred_assets( + tree: DocTreeNode, + pending_assets: list[dict], +) -> None: + """Place collected assets into the tree based on final navigated paths. + + Called ONCE after the entire BFS + discovery merge completes for a + document. Only assets whose ``owner_section_path`` exactly matches + a final leaf_content key are retained — all others are discarded. + This implements the "subset refinement" strategy: assets fetched at + a broad scope are filtered down to the paths the LLM actually + selected across all BFS depths. + """ + final_paths = _collect_all_leaf_paths(tree) + if not final_paths: + return + + # Collect existing chunk_ids to avoid duplicates + existing_ids = { + str(row.get('chunk_id') or '') + for row in tree.flatten_chunk_rows() + if row.get('chunk_id') + } + + placed = 0 + for asset in pending_assets: + chunk_id = str(asset.get('chunk_id') or '') + if chunk_id and chunk_id in existing_ids: + continue # already in tree via hydrate_connected_target_rows + + owner_path = ( + asset.get('owner_section_path') + or asset.get('section_path') + ) + if not owner_path or owner_path not in final_paths: + continue # owner not in any navigated path → discard + + # Place into root; reparent_leaf_content will move to correct child + tree.add_leaf_chunks(owner_path, [asset]) + if chunk_id: + existing_ids.add(chunk_id) + placed += 1 + + if placed: + tree.reparent_leaf_content() + logger.info( + f' deferred asset reconcile: {placed}/{len(pending_assets)} ' + f'assets placed into {len(final_paths)} navigated paths' + ) + + def _build_config_from_env() -> AgentRunConfig: """Read agent config from environment, with sensible defaults.""" return AgentRunConfig( @@ -412,6 +471,9 @@ async def run( channels: list[str] | None = None, channel_weights: dict[str, float] | None = None, config: AgentRunConfig | None = None, + ledger: BudgetLedger | None = None, + parent_run_id: str | None = None, + workflow_step_id: str | None = None, ) -> AgenticResult: """Run the agentic retrieval pipeline. @@ -431,7 +493,7 @@ async def run( exclude_sections = exclude_sections or [] state = AgentState() - state.ledger = BudgetLedger( + state.ledger = ledger or BudgetLedger( total=config.token_budget_total, planning_ratio=config.planning_ratio, bootstrap=config.bootstrap_budget, @@ -455,6 +517,8 @@ async def run( 'exclude_sections': exclude_sections, 'signal_paths': signal_paths, }, + parent_run_id=parent_run_id, + workflow_step_id=workflow_step_id, ) trace_enabled = os.environ.get('RETRIEVAL_AGENTIC_TRACE_ENABLED', 'true') == 'true' @@ -714,6 +778,7 @@ async def _context_llm_call(prompt): # BFS queue: (scope_path, parent_node, depth) root = DocTreeNode(scope_path=None) pending: list[tuple[str | None, DocTreeNode, int]] = [(None, root, 0)] + doc_pending_assets: list[dict] = [] # deferred asset reconcile while pending: if state.elapsed_ms >= config.latency_budget_ms: @@ -777,11 +842,10 @@ async def _context_llm_call(prompt): f'depth={depth} tools={tool_choices or ["NAVIGATE"]}' ) - pending_scope_assets: list[dict] = [] for asset_tool in tool_choices: if asset_tool not in ('FIND_IMAGES', 'FIND_TABLES'): continue - # ★ Asset collection (programmatic extraction) + # ★ Asset collection → deferred reconcile asset_type = 'image' if asset_tool == 'FIND_IMAGES' else 'table' asset_chunks = await tools.asset_filter_step( db, @@ -791,7 +855,7 @@ async def _context_llm_call(prompt): asset_type=asset_type, ) if asset_chunks: - pending_scope_assets.extend(asset_chunks) + doc_pending_assets.extend(asset_chunks) if trace_enabled: trace.record_step( @@ -843,34 +907,7 @@ async def _context_llm_call(prompt): parent_node.add_leaf_chunks(leaf_path, chunks) parent_node.confidence = step_node.confidence - # ★ Step 2.5: Reconcile pending assets with navigated leaf content - if pending_scope_assets: - # Collect all chunk_ids already present in any leaf_content - existing_ids = { - str(row.get('chunk_id') or '') - for row in parent_node.flatten_chunk_rows() - if row.get('chunk_id') - } - # Filter out assets already inlined - supplementary = [ - a for a in pending_scope_assets - if str(a.get('chunk_id') or '') not in existing_ids - ] - if supplementary: - # Only place assets into sections already in the - # navigated tree. If the owner path isn't part of - # the tree, fall back to the current scope. - navigated_paths = set(parent_node.leaf_content.keys()) | set(parent_node.children.keys()) - for asset in supplementary: - owner_path = ( - asset.get('owner_section_path') - or asset.get('section_path') - or scope - ) - if owner_path and owner_path in navigated_paths: - parent_node.add_leaf_chunks(str(owner_path), [asset]) - elif scope: - parent_node.add_leaf_chunks(str(scope), [asset]) + # (Asset reconcile is deferred until after BFS + discovery merge) # Accumulate hydrated leaf paths into doc_exclude # so subsequent drill-downs don't re-show them as [SELECT]. @@ -975,6 +1012,32 @@ async def _context_llm_call(prompt): chunks=sum(len(chunks) for chunks in discovery_node.leaf_content.values()), ) + # ── Deferred asset reconcile ────────────────────────────── + # Assets were collected across all BFS depths but NOT placed + # into the tree yet. Now that the final navigated paths are + # known (BFS + discovery), filter and place only those assets + # whose owner path matches a navigated leaf. + if not is_b_class and doc_pending_assets: + _reconcile_deferred_assets(root, doc_pending_assets) + if trace_enabled: + trace.record_step( + 'deferred_asset_reconcile', ToolResult( + status='reconciled', + payload={ + 'document_id': doc.document_id, + 'pending_count': len(doc_pending_assets), + 'placed_count': sum( + 1 for a in doc_pending_assets + if str(a.get('chunk_id') or '') in { + str(r.get('chunk_id') or '') + for r in root.flatten_chunk_rows() + } + ), + }, + ), + decision_reason=f'deferred_reconcile_r{round_idx}_{doc.source_file_name}', + ) + # Merge or store doc tree if doc.document_id in state.doc_trees: state.doc_trees[doc.document_id].merge(root) diff --git a/packages/shared-python/shared/services/retrieval/agentic/tools.py b/packages/shared-python/shared/services/retrieval/agentic/tools.py index d11e4da7b..588e51dd5 100644 --- a/packages/shared-python/shared/services/retrieval/agentic/tools.py +++ b/packages/shared-python/shared/services/retrieval/agentic/tools.py @@ -381,9 +381,10 @@ async def kg_document_select( === Available Actions === -NAVIGATE (always performed) - Drill into specific sections to explore detailed content. - This action always runs — you do not need to select it. +NAVIGATE (separate step) + A separate navigation decision follows this step. The LLM may choose + specific sections to drill into, or decide not to drill deeper. + You do NOT need to select NAVIGATE here — it is handled separately. FIND_IMAGES (optional, additive) Also extract image/chart/diagram assets under this scope. @@ -394,7 +395,7 @@ async def kg_document_select( Select this when the query asks about tables, tabular data, or structured data. You may select ZERO, ONE, or BOTH optional actions. -Navigation always happens regardless of your selection. +Navigation is decided separately and does not conflict with these actions. Return ONLY a JSON object: {{"tools": []}} — navigate only, no extra assets @@ -927,13 +928,15 @@ async def discovery_select_step( t0 = time.monotonic() try: - # 1. Format hints for LLM + # 1. Format hints for LLM (deduplicate by section_path) hint_lines: list[str] = [] hint_by_path: dict[str, dict] = {} for h in hints: sp = h.get('section_path', '') if not sp or sp == 'Root': continue + if sp in hint_by_path: + continue # skip duplicate section_path title = sp.rsplit(' / ', 1)[-1] if ' / ' in sp else sp summary = h.get('summary', '') or '' hint_lines.append(f'▸ path="{sp}" {title} [Leaf]') diff --git a/packages/shared-python/shared/services/retrieval/agentic/trace.py b/packages/shared-python/shared/services/retrieval/agentic/trace.py index 8a58a243c..9c48e7a30 100644 --- a/packages/shared-python/shared/services/retrieval/agentic/trace.py +++ b/packages/shared-python/shared/services/retrieval/agentic/trace.py @@ -50,6 +50,9 @@ def __init__( top_k: int = 10, data_type: int = 1, filters: dict[str, Any] | None = None, + parent_run_id: str | None = None, + workflow_step_id: str | None = None, + workflow_plan: dict[str, Any] | None = None, ) -> None: self._db = db self._run_id = f'aret_{uuid4().hex[:12]}' @@ -60,6 +63,9 @@ def __init__( self._top_k = top_k self._data_type = data_type self._filters = filters or {} + self._parent_run_id = parent_run_id + self._workflow_step_id = workflow_step_id + self._workflow_plan = workflow_plan self._steps: list[dict[str, Any]] = [] self._start_time = time.monotonic() self._created = False @@ -86,6 +92,9 @@ async def create_run(self) -> None: agentic_enabled=True, cache_hit=False, result_count=0, + parent_run_id=self._parent_run_id, + workflow_step_id=self._workflow_step_id, + workflow_plan=self._workflow_plan, latency_ms=0, created_at=_now_utc(), ) @@ -179,6 +188,10 @@ async def complete( } if budget_snapshot is not None: provenance['budget_snapshot'] = budget_snapshot + if self._parent_run_id: + provenance['parent_run_id'] = self._parent_run_id + if self._workflow_step_id: + provenance['workflow_step_id'] = self._workflow_step_id stmt = ( update(RetrievalRun) diff --git a/packages/shared-python/shared/services/retrieval/app_service.py b/packages/shared-python/shared/services/retrieval/app_service.py index c066a1c16..dc14c055a 100644 --- a/packages/shared-python/shared/services/retrieval/app_service.py +++ b/packages/shared-python/shared/services/retrieval/app_service.py @@ -982,6 +982,7 @@ async def run_retrieval_query( rerank: bool = False, threshold: float = 0.0, internal_recall_k: int | None = None, + enable_decomposition: bool | None = None, # deprecated: now always uses workflow ) -> dict[str, Any]: """Checkerboard retrieval: 3 independent channels -> RRF -> agent/graph union -> assembly.""" t_start = time.monotonic() @@ -1015,6 +1016,8 @@ async def run_retrieval_query( rerank=rerank, threshold=threshold, internal_recall_k=internal_recall_k, + # Always True: agentic mode now always routes through workflow + decomposition_enabled=True, ) cache_version: int | None = None @@ -1096,22 +1099,22 @@ async def run_retrieval_query( logger.info(f' ✅ Small KB: {len(results)} results in {elapsed_total}ms') return await _to_public_response(response) - # ══ Route: agentic vs legacy ══ + # ══ Route: agentic (unified workflow) vs legacy ══ _agentic_enabled = os.environ.get('RETRIEVAL_AGENTIC_ENABLED', 'false') == 'true' if _agentic_enabled: - # ── AGENTIC path (all errors self-contained, no fallback to legacy) ── - from shared.services.retrieval.agentic.orchestrator import RetrievalAgent - from shared.services.retrieval.llm_adapter import create_retrieval_llm_fn as _create_llm - - llm_fn = _create_llm() - agent = RetrievalAgent() - agentic_result = await agent.run( + # ── Unified agentic path via WorkflowOrchestrator ── + # Simple queries: planner returns a single-step plan (no decomposition). + # Complex queries: planner returns a multi-step plan with synthesize. + # Both go through the same code path. + from shared.services.retrieval.workflow.orchestrator import WorkflowOrchestrator + + workflow = WorkflowOrchestrator() + workflow_result = await workflow.run( db, user_id=user_id, namespace=namespace, query=query, top_k=top_k, - llm_fn=llm_fn, exclude_document_ids=exclude_document_ids, exclude_sections=exclude_sections, data_type=data_type, @@ -1120,11 +1123,10 @@ async def run_retrieval_query( channels=channels, channel_weights=channel_weights, ) - router_used = agentic_result.router_used - # Generate asset URLs for media chunks in referenced_chunks + # Enrich referenced_chunks with asset URLs (images/tables) enriched_refs: list[dict[str, Any]] = [] - for ref in agentic_result.referenced_chunks: + for ref in workflow_result.referenced_chunks: enriched = dict(ref) chunk_type = _normalize_chunk_type(ref.get('chunk_type')) artifact_ref = ref.get('file_path', '') @@ -1140,30 +1142,9 @@ async def run_retrieval_query( logger.warning(f'Failed to generate agentic asset URL (ignored): {e}') enriched_refs.append(enriched) - # Build backward-compatible results[] from referenced_chunks - # (minimal: chunk_id + document_id + chunk_type + section_path) - results = [ - { - 'chunk_id': ref.get('chunk_id'), - 'document_id': ref.get('document_id'), - 'chunk_type': ref.get('chunk_type'), - 'source': { - 'document_id': ref.get('document_id'), - 'section_path': ref.get('section_path'), - }, - } - for ref in enriched_refs - ] - - response = { - "namespace": namespace, - "query": query, - "router_used": router_used, - "results": results, - "evidence_text": agentic_result.evidence_text, - "answer_text": agentic_result.answer_text, - "referenced_chunks": enriched_refs, - } + response = workflow_result.to_api_response() + # Override referenced_chunks with enriched versions + response['referenced_chunks'] = enriched_refs if cache_version is not None: try: @@ -1180,7 +1161,7 @@ async def run_retrieval_query( try: schedule_retrieval_hit_stats_update( user_id=user_id, namespace=namespace, - results=agentic_result.referenced_chunks, + results=enriched_refs, ) except Exception as e: logger.warning(f"Failed to trigger retrieval hit stats update (ignored): {e}") @@ -1190,14 +1171,15 @@ async def run_retrieval_query( f'\n{"█" * 70}\n' f' ✅ AGENTIC RETRIEVAL COMPLETE: ' f'{len(enriched_refs)} chunks | ' - f'evidence={len(agentic_result.evidence_text)} chars | ' - f'answer={len(agentic_result.answer_text)} chars | ' - f'router={router_used} | {elapsed_total}ms\n' + f'answer={len(workflow_result.answer_text)} chars | ' + f'router={workflow_result.router_used} | {elapsed_total}ms\n' f'{"█" * 70}' ) return await _to_public_response(response) + else: + # ── LEGACY path (existing code, unchanged) ── # ── Channel execution ── diff --git a/packages/shared-python/shared/services/retrieval/cache_service.py b/packages/shared-python/shared/services/retrieval/cache_service.py index b7c2bd88b..bb165eb41 100644 --- a/packages/shared-python/shared/services/retrieval/cache_service.py +++ b/packages/shared-python/shared/services/retrieval/cache_service.py @@ -8,6 +8,7 @@ from shared.services.redis import RedisServiceFactory _RETRIEVAL_CACHE_TTL_SECONDS = 300 +_WORKFLOW_PLAN_CACHE_TTL_SECONDS = 600 _VERSION_FALLBACK = 0 @@ -42,6 +43,7 @@ def _cache_shape_digest( rerank: bool = False, threshold: float = 0.0, internal_recall_k: int | None = None, + decomposition_enabled: bool | None = None, ) -> str: normalized_excludes = sorted(exclude_document_ids) normalized_sections = _normalize_exclude_sections(exclude_sections) @@ -55,6 +57,7 @@ def _cache_shape_digest( str(rerank), str(threshold), str(internal_recall_k), + str(decomposition_enabled), ] ) payload = f"{query}|{top_k}|{'|'.join(normalized_excludes)}|{'|'.join(normalized_sections)}|{extra}" @@ -178,3 +181,37 @@ async def set_cached_retrieval_query_result( response, ex=_RETRIEVAL_CACHE_TTL_SECONDS, ) + + +def _workflow_plan_cache_key(*, user_id: str, namespace: str, query: str) -> str: + digest = hashlib.sha256(query.encode("utf-8")).hexdigest() + return f"retrieval:workflow:plan:{user_id}:{namespace}:{digest}" + + +async def get_cached_workflow_plan( + *, + user_id: str, + namespace: str, + query: str, +) -> dict[str, Any] | None: + redis_service = RedisServiceFactory.get_service() + cached = await redis_service.get( + _workflow_plan_cache_key(user_id=user_id, namespace=namespace, query=query), + default=None, + ) + return cached if isinstance(cached, dict) else None + + +async def set_cached_workflow_plan( + *, + user_id: str, + namespace: str, + query: str, + plan: dict[str, Any], +) -> None: + redis_service = RedisServiceFactory.get_service() + await redis_service.set( + _workflow_plan_cache_key(user_id=user_id, namespace=namespace, query=query), + plan, + ex=_WORKFLOW_PLAN_CACHE_TTL_SECONDS, + ) diff --git a/packages/shared-python/shared/services/retrieval/llm_adapter.py b/packages/shared-python/shared/services/retrieval/llm_adapter.py index c6566d5f3..c5e592b5e 100644 --- a/packages/shared-python/shared/services/retrieval/llm_adapter.py +++ b/packages/shared-python/shared/services/retrieval/llm_adapter.py @@ -54,6 +54,21 @@ def _resolve_default_model() -> str: return getattr(settings, 'NORMOL_MODEL', None) or 'deepseek-chat' +def _resolve_planner_model(*, thinking: bool) -> str: + configured = getattr(settings, 'RETRIEVAL_PLANNER_MODEL', '') or '' + if configured: + return configured + if getattr(settings, 'DS_KEY', ''): + return 'deepseek-reasoner' if thinking else 'deepseek-chat' + if getattr(settings, 'ALI_API_KEYS', ''): + return 'qwq-32b-preview' if thinking else 'qwen-plus' + if getattr(settings, 'GLM_API_KEY', ''): + return 'glm-4-plus' if thinking else 'glm-4-flash' + if getattr(settings, 'GPT_API_KEY', ''): + return 'o3-mini' if thinking else 'gpt-4o-mini' + return getattr(settings, 'NORMOL_MODEL', None) or 'deepseek-chat' + + def create_retrieval_llm_fn( *, model: str | None = None, @@ -89,6 +104,37 @@ async def llm_fn(prompt: LLMFnInput) -> str: return llm_fn +def create_retrieval_planner_fn( + *, + thinking: bool = True, + model: str | None = None, + max_tokens: int = 8192, +) -> LLMFn | None: + """Create a reasoning-capable LLM callable for query planning.""" + if not _has_llm_credentials(): + logger.debug('retrieval: no LLM credentials configured, workflow planner disabled') + return None + + effective_model = model or _resolve_planner_model(thinking=thinking) + + async def llm_fn(prompt: LLMFnInput) -> str: + from shared.utils.OpenAICompatibleClientSync import get_openai_client + + client = get_openai_client(model=effective_model) + current_llm_usage.set(None) + result, usage = await asyncio.to_thread( + client.chat_completion_with_usage, + cast(Any, prompt), + model=effective_model, + temperature=0.0, + max_tokens=max_tokens, + ) + current_llm_usage.set(usage) + return result + + return llm_fn + + def create_retrieval_vlm_fn( *, model: str | None = None, diff --git a/packages/shared-python/shared/services/retrieval/workflow/__init__.py b/packages/shared-python/shared/services/retrieval/workflow/__init__.py new file mode 100644 index 000000000..67f896d3d --- /dev/null +++ b/packages/shared-python/shared/services/retrieval/workflow/__init__.py @@ -0,0 +1,13 @@ +"""Query-decomposition retrieval workflow.""" +from .types import PlannedStep, QueryPlan, StepResult, WorkflowResult +from .wallet import BudgetWallet +from .orchestrator import WorkflowOrchestrator + +__all__ = [ + "BudgetWallet", + "PlannedStep", + "QueryPlan", + "StepResult", + "WorkflowResult", + "WorkflowOrchestrator", +] diff --git a/packages/shared-python/shared/services/retrieval/workflow/orchestrator.py b/packages/shared-python/shared/services/retrieval/workflow/orchestrator.py new file mode 100644 index 000000000..9c5a20aff --- /dev/null +++ b/packages/shared-python/shared/services/retrieval/workflow/orchestrator.py @@ -0,0 +1,382 @@ +"""Workflow orchestrator for decomposed retrieval queries.""" +from __future__ import annotations + +import asyncio +import os +import time +from typing import Any +from uuid import uuid4 + +from loguru import logger +from sqlalchemy.ext.asyncio import AsyncSession + +from shared.core.database import get_db_context +from shared.services.retrieval.agentic.budget import BudgetLedger +from shared.services.retrieval.agentic.orchestrator import RetrievalAgent +from shared.services.retrieval.agentic.types import AgenticResult +from shared.services.retrieval.cache_service import ( + get_cached_workflow_plan, + set_cached_workflow_plan, +) +from shared.services.retrieval.llm_adapter import ( + create_retrieval_llm_fn, + create_retrieval_planner_fn, +) +from shared.services.retrieval.workflow.planner import QueryPlanner +from shared.services.retrieval.workflow.synthesizer import compose_final_answer, synthesize_step +from shared.services.retrieval.workflow.types import PlannedStep, QueryPlan, StepResult, WorkflowResult +from shared.services.retrieval.workflow.wallet import BudgetWallet + + +class WorkflowOrchestrator: + """Plan and execute a query workflow DAG.""" + + def __init__(self) -> None: + self.parent_run_id = f'wret_{uuid4().hex[:12]}' + + async def run( + self, + db: AsyncSession, + *, + user_id: str, + namespace: str, + query: str, + top_k: int, + exclude_document_ids: list[str], + exclude_sections: list[dict[str, str]], + data_type: int = 1, + signal_paths: list[str] | None = None, + filter_mode: str = 'delete', + channels: list[str] | None = None, + channel_weights: dict[str, float] | None = None, + llm_fn=None, + ) -> WorkflowResult: + t0 = time.monotonic() + llm_fn = llm_fn or create_retrieval_llm_fn() + planner_llm = create_retrieval_planner_fn(thinking=True) + planner_budget = _env_int('RETRIEVAL_PLANNER_THINKING_BUDGET', 4000) + wallet_total = _env_int('RETRIEVAL_WALLET_TOTAL_BUDGET', 200000) + per_retrieve = _env_int('RETRIEVAL_WALLET_PER_RETRIEVE_STEP_BUDGET', 40000) + per_synthesize = _env_int('RETRIEVAL_WALLET_PER_SYNTHESIZE_STEP_BUDGET', 6000) + max_steps = _env_int('RETRIEVAL_DECOMPOSITION_MAX_STEPS', 5) + + planner_ledger = BudgetLedger( + total=planner_budget, + planning_ratio=0.0, + bootstrap=planner_budget, + per_doc_min_share=0, + ) + plan = await self._load_or_plan( + user_id=user_id, + namespace=namespace, + query=query, + planner_llm=planner_llm, + planner_ledger=planner_ledger, + max_steps=max_steps, + wallet_total=wallet_total, + per_retrieve=per_retrieve, + ) + + wallet = BudgetWallet( + total=wallet_total, + per_retrieve_step_default=per_retrieve, + per_synthesize_step_default=per_synthesize, + ) + ledgers = await wallet.allocate(plan) + results_by_id: dict[str, StepResult] = {} + sem = asyncio.Semaphore(_env_int('RETRIEVAL_WORKFLOW_PARALLEL_MAX', 3)) + + for batch in plan.topological_batches(): + await asyncio.gather( + *[ + self._run_step( + db, + step=step, + ledger=ledgers[step.id], + results_by_id=results_by_id, + semaphore=sem, + user_id=user_id, + namespace=namespace, + top_k=step.top_k or top_k, + exclude_document_ids=exclude_document_ids, + exclude_sections=exclude_sections, + data_type=step.data_type or data_type, + signal_paths=signal_paths, + filter_mode=filter_mode, + channels=channels, + channel_weights=channel_weights, + llm_fn=llm_fn, + ) + for step in batch + ] + ) + for step in batch: + await wallet.reclaim(step.id, ledgers[step.id]) + + answer_text = compose_final_answer(plan, results_by_id) + ordered_results = [results_by_id[step.id] for step in plan.steps if step.id in results_by_id] + referenced_chunks = _dedupe_references( + ref for step_result in ordered_results for ref in step_result.referenced_chunks + ) + api_results = _references_to_results(referenced_chunks) + elapsed_ms = int((time.monotonic() - t0) * 1000) + logger.info( + 'workflow retrieval DONE: steps={} refs={} answer_chars={} elapsed={}ms', + len(ordered_results), + len(referenced_chunks), + len(answer_text), + elapsed_ms, + ) + return WorkflowResult( + namespace=namespace, + query=query, + router_used='workflow_decomposed' if len(plan.steps) > 1 else 'workflow_single_step', + answer_text=answer_text, + plan=plan, + steps=ordered_results, + referenced_chunks=referenced_chunks, + results=api_results, + final_strategy_used=plan.final_strategy, + wallet_snapshot=wallet.snapshot(), + planner_snapshot=planner_ledger.snapshot(), + parent_run_id=self.parent_run_id, + ) + + async def _load_or_plan( + self, + *, + user_id: str, + namespace: str, + query: str, + planner_llm, + planner_ledger: BudgetLedger, + max_steps: int, + wallet_total: int, + per_retrieve: int, + ) -> QueryPlan: + try: + cached = await get_cached_workflow_plan(user_id=user_id, namespace=namespace, query=query) + if cached: + return QueryPlan.from_dict(cached, original_query=query) + except Exception as exc: + logger.warning(f'workflow plan cache read failed (ignored): {exc}') + + planner = QueryPlanner( + llm_fn=planner_llm, + planner_ledger=planner_ledger, + max_steps=max_steps, + total_budget=wallet_total, + per_step_budget=per_retrieve, + ) + plan = await planner.plan(query=query) + try: + await set_cached_workflow_plan( + user_id=user_id, + namespace=namespace, + query=query, + plan=plan.to_dict(), + ) + except Exception as exc: + logger.warning(f'workflow plan cache write failed (ignored): {exc}') + return plan + + async def _run_step( + self, + db: AsyncSession, + *, + step: PlannedStep, + ledger: BudgetLedger, + results_by_id: dict[str, StepResult], + semaphore: asyncio.Semaphore, + user_id: str, + namespace: str, + top_k: int, + exclude_document_ids: list[str], + exclude_sections: list[dict[str, str]], + data_type: int, + signal_paths: list[str] | None, + filter_mode: str, + channels: list[str] | None, + channel_weights: dict[str, float] | None, + llm_fn, + ) -> None: + async with semaphore: + if step.step_kind == 'synthesize': + await self._run_synthesize_step(step, ledger, results_by_id, llm_fn) + return + await self._run_retrieve_step( + db, + step=step, + ledger=ledger, + results_by_id=results_by_id, + user_id=user_id, + namespace=namespace, + top_k=top_k, + exclude_document_ids=exclude_document_ids, + exclude_sections=exclude_sections, + data_type=data_type, + signal_paths=signal_paths, + filter_mode=filter_mode, + channels=channels, + channel_weights=channel_weights, + llm_fn=llm_fn, + ) + + async def _run_retrieve_step( + self, + db: AsyncSession, + *, + step: PlannedStep, + ledger: BudgetLedger, + results_by_id: dict[str, StepResult], + user_id: str, + namespace: str, + top_k: int, + exclude_document_ids: list[str], + exclude_sections: list[dict[str, str]], + data_type: int, + signal_paths: list[str] | None, + filter_mode: str, + channels: list[str] | None, + channel_weights: dict[str, float] | None, + llm_fn, + ) -> None: + try: + # AsyncSession is not safe for concurrent use. Workflow steps may + # run in the same topological batch, so each retrieve step opens an + # isolated session and leaves the parent session untouched. + async with get_db_context() as step_db: + agentic_result = await RetrievalAgent().run( + step_db, + user_id=user_id, + namespace=namespace, + query=step.sub_query, + top_k=top_k, + llm_fn=llm_fn, + exclude_document_ids=exclude_document_ids, + exclude_sections=exclude_sections, + data_type=data_type, + signal_paths=signal_paths, + filter_mode=filter_mode, + channels=channels, + channel_weights=channel_weights, + ledger=ledger, + parent_run_id=self.parent_run_id, + workflow_step_id=step.id, + ) + results_by_id[step.id] = _step_result_from_agentic(step, agentic_result) + except Exception as exc: + logger.exception(f'workflow retrieve step failed: step_id={step.id}') + results_by_id[step.id] = StepResult( + step_id=step.id, + sub_query=step.sub_query, + step_kind=step.step_kind, + depends_on=step.depends_on, + output_role=step.output_role, + status='error', + error=str(exc), + budget_snapshot=ledger.snapshot(), + ) + + async def _run_synthesize_step( + self, + step: PlannedStep, + ledger: BudgetLedger, + results_by_id: dict[str, StepResult], + llm_fn, + ) -> None: + if llm_fn is None: + results_by_id[step.id] = StepResult( + step_id=step.id, + sub_query=step.sub_query, + step_kind=step.step_kind, + depends_on=step.depends_on, + output_role=step.output_role, + status='skipped', + answer_text='', + error='llm unavailable for synthesis', + budget_snapshot=ledger.snapshot(), + ) + return + prior = {dep: results_by_id[dep] for dep in step.depends_on if dep in results_by_id} + try: + answer = await synthesize_step(step, prior_results=prior, llm_fn=llm_fn, ledger=ledger) + refs = _dedupe_references( + ref for result in prior.values() for ref in result.referenced_chunks + ) + results_by_id[step.id] = StepResult( + step_id=step.id, + sub_query=step.sub_query, + step_kind=step.step_kind, + depends_on=step.depends_on, + output_role=step.output_role, + status='done', + answer_text=answer, + referenced_chunks=refs, + budget_snapshot=ledger.snapshot(), + ) + except Exception as exc: + results_by_id[step.id] = StepResult( + step_id=step.id, + sub_query=step.sub_query, + step_kind=step.step_kind, + depends_on=step.depends_on, + output_role=step.output_role, + status='budget_stop' if 'budget' in str(exc).lower() else 'error', + answer_text='(budget exhausted)' if 'budget' in str(exc).lower() else '', + error=str(exc), + budget_snapshot=ledger.snapshot(), + ) + + +def _step_result_from_agentic(step: PlannedStep, result: AgenticResult) -> StepResult: + status = 'budget_stop' if 'budget' in (result.stop_reason or '') else 'done' + return StepResult( + step_id=step.id, + sub_query=step.sub_query, + step_kind=step.step_kind, + depends_on=step.depends_on, + output_role=step.output_role, + status=status, # type: ignore[arg-type] + answer_text=result.answer_text, + evidence_text=result.evidence_text, + referenced_chunks=result.referenced_chunks, + budget_snapshot=result.budget_snapshot, + router_used=result.router_used, + stop_reason=result.stop_reason, + ) + + +def _dedupe_references(refs) -> list[dict[str, Any]]: + seen: set[str] = set() + out: list[dict[str, Any]] = [] + for ref in refs: + chunk_id = str(ref.get('chunk_id') or '') + key = chunk_id or str(ref) + if key in seen: + continue + seen.add(key) + out.append(dict(ref)) + return out + + +def _references_to_results(refs: list[dict[str, Any]]) -> list[dict[str, Any]]: + return [ + { + 'chunk_id': ref.get('chunk_id'), + 'document_id': ref.get('document_id'), + 'chunk_type': ref.get('chunk_type'), + 'source': { + 'document_id': ref.get('document_id'), + 'section_path': ref.get('section_path'), + }, + } + for ref in refs + ] + + +def _env_int(name: str, default: int) -> int: + try: + return int(os.environ.get(name, str(default))) + except (TypeError, ValueError): + return default diff --git a/packages/shared-python/shared/services/retrieval/workflow/planner.py b/packages/shared-python/shared/services/retrieval/workflow/planner.py new file mode 100644 index 000000000..e34713db6 --- /dev/null +++ b/packages/shared-python/shared/services/retrieval/workflow/planner.py @@ -0,0 +1,239 @@ +"""Query planner for decomposed retrieval workflows.""" +from __future__ import annotations + +import json +import re +import time +from typing import Any + +from loguru import logger + +from shared.services.retrieval.agentic.budget import BudgetExceeded, BudgetLedger +from shared.services.retrieval.llm_adapter import LLMFn, current_llm_usage +from shared.services.retrieval.workflow.types import FinalStrategy, OutputRole, PlannedStep, QueryPlan, StepKind +from shared.utils.token_estimate import estimate_tokens + + +_PLAN_SCHEMA = { + "reasoning_summary": "