import logging from functools import lru_cache from typing import Any, cast from sqlalchemy.ext.asyncio import AsyncSession from app.core.config import get_settings from app.core.errors import RecoverableAgentError from app.core.fund_contracts import FundQuoteQuery from app.core.nl2sql_contracts import FinancialNL2SQLInput from app.infrastructure.fund_quote_cache import FundQuoteCache from app.infrastructure.memory_cache import MemoryCacheAdapter from app.infrastructure.vector_memory import VectorMemoryAdapter from app.service.agent.factory import AgentFactory from app.service.agent.governance import PlatformGovernance from app.service.agent.offsite_fund_agent import OffsiteFundAgent from app.service.agent.promotion_material_agent import PromotionMaterialAgent from app.service.financial_nl2sql_service import query_financial_data_tool from app.service.fund_quote_service import query_fund_quote_tool from app.service.intent_classifier import IntentClassifier from app.service.memory_recall_service import MemoryRecallService from app.service.model_gateway import ( DatabaseModelEndpointResolver, DatabaseModelGateway, ModelDispatchService, ModelEmbeddingService, ModelGenerationService, ) from app.service.runtime_config_service import load_active_intent_configs from app.service.suitability_service import SuitabilityToolInput, suitability_tool_handler from app.service.tool_executor import ToolDefinition, ToolExecutor, ToolRegistry logger = logging.getLogger(__name__) @lru_cache(maxsize=1) def get_model_service() -> ModelGenerationService: """模型生成的统一装配入口,HTTP 与 Worker 共用。""" return ModelGenerationService(ModelDispatchService(DatabaseModelGateway())) @lru_cache(maxsize=1) def get_memory_cache_adapter() -> MemoryCacheAdapter | None: """Redis 不可用时降级为无缓存,不能阻塞业务调用。""" try: from redis.asyncio import Redis settings = get_settings() client = Redis.from_url( settings.redis_url, socket_connect_timeout=settings.redis_connect_timeout_seconds, socket_timeout=settings.redis_connect_timeout_seconds, decode_responses=True, ) return MemoryCacheAdapter(client) except Exception: logger.warning("memory cache adapter unavailable; recall runs without cache", exc_info=True) return None @lru_cache(maxsize=1) def get_fund_quote_cache() -> FundQuoteCache | None: """行情缓存与记忆缓存共用 Redis 连接配置,故障时透明降级。""" adapter = get_memory_cache_adapter() return FundQuoteCache(adapter) if adapter is not None else None @lru_cache(maxsize=1) def get_vector_memory_adapter() -> VectorMemoryAdapter | None: """Milvus 语义召回适配器不可用时关闭语义通道,保留结构化召回。""" try: from pymilvus import MilvusClient # type: ignore[import-untyped] settings = get_settings() client = MilvusClient(uri=settings.milvus_uri, token=settings.milvus_token or None) return VectorMemoryAdapter(client, settings.milvus_collection) except Exception: logger.warning("vector memory adapter unavailable; semantic recall disabled", exc_info=True) return None @lru_cache(maxsize=1) def get_memory_embedding_service() -> ModelEmbeddingService: return ModelEmbeddingService(ModelDispatchService(DatabaseModelGateway())) async def _embed_text(text: str) -> list[float]: endpoints = await DatabaseModelEndpointResolver().resolve( agent_type="memory_recall", task_type="embedding" ) if not endpoints: raise RecoverableAgentError("没有可用的 embedding 端点") return (await get_memory_embedding_service().embed(endpoints, text)).vector def build_memory_recall_service(session: AsyncSession) -> MemoryRecallService: vector = get_vector_memory_adapter() return MemoryRecallService( session, vector=vector, cache=get_memory_cache_adapter(), embed=_embed_text if vector is not None else None, ) @lru_cache(maxsize=1) def get_agent_factory() -> AgentFactory: """HTTP 与 Worker 共用的唯一底座依赖组装入口。""" registry = ToolRegistry() registry.register(ToolDefinition( name="check_suitability", input_model=SuitabilityToolInput, handler=cast(Any, suitability_tool_handler), required_permission="suitability:read", allowed_roles=("customer", "advisor", "operator", "admin"), )) registry.register(ToolDefinition( name="query_fund_quote", input_model=FundQuoteQuery, handler=cast(Any, query_fund_quote_tool), required_permission="fund:quote:read", allowed_roles=("customer", "advisor", "operator", "risk_operator", "admin"), timeout_seconds=15, )) registry.register(ToolDefinition( name="query_financial_data", input_model=FinancialNL2SQLInput, handler=cast(Any, query_financial_data_tool), required_permission="financial:nl2sql:read", allowed_roles=("advisor", "operator", "admin", "super_admin"), timeout_seconds=10, )) model_service = get_model_service() endpoint_resolver = DatabaseModelEndpointResolver() factory = AgentFactory( governance=PlatformGovernance(recall_factory=build_memory_recall_service), model_service=model_service, tool_executor=ToolExecutor(registry), intent_classifier=IntentClassifier( model_service, config_loader=load_active_intent_configs ), intent_endpoint_resolver=endpoint_resolver, ) register_business_agents(factory) return factory def register_business_agents(factory: AgentFactory) -> None: """业务 Agent 仍由本项目统一注册,基座只负责提供公共治理链。""" factory.register( OffsiteFundAgent.definition, lambda _context: OffsiteFundAgent(OffsiteFundAgent.definition), ) factory.register( PromotionMaterialAgent.definition, lambda _context: PromotionMaterialAgent(PromotionMaterialAgent.definition), )