161 lines
6.1 KiB
Python
161 lines
6.1 KiB
Python
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),
|
|
)
|