From 5915319f9fa805fc91b495227485fc1727a0a4cd Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E5=8F=B2=E8=8D=92=E4=B8=98?= Date: Fri, 11 Sep 2026 10:47:01 +0800 Subject: [PATCH 1/3] =?UTF-8?q?fix:=20=E6=9B=B4=E6=96=B0main.py=E9=80=BB?= =?UTF-8?q?=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/deps.py | 14 +++++++++++++- api/router.py | 5 +++++ config/database/milvus.py | 12 +++++++++++- config/settings.py | 9 +++++++-- main.py | 12 +++++++++++- nl2sql/__init__.py | 0 requirements.txt | 15 +++++++++++++++ sql/schema.sql | 5 +++-- utils/response.py | 2 +- 9 files changed, 66 insertions(+), 8 deletions(-) create mode 100644 nl2sql/__init__.py diff --git a/api/deps.py b/api/deps.py index 7d75ee1..94deffe 100644 --- a/api/deps.py +++ b/api/deps.py @@ -9,6 +9,9 @@ from service.auth import decode_token from utils.exceptions import AuthError, ForbiddenError +KNOWLEDGE_OPERATOR_ROLES = {"ADMIN", "KNOWLEDGE_ADMIN", "KNOWLEDGE_OPERATOR", "运营"} + + async def get_current_user( request: Request, db: AsyncSession = Depends(get_db) ) -> SysUser: @@ -24,4 +27,13 @@ async def get_current_user( raise AuthError("用户不存在") if user.status != "正常": raise ForbiddenError("账号已被禁用或冻结") - return user \ No newline at end of file + return user + + +async def require_knowledge_operator(user: SysUser = Depends(get_current_user)) -> SysUser: + """Allow only administrators or explicitly assigned knowledge operators.""" + if user.user_type == "ADMIN": + return user + if user.user_type != "EMPLOYEE" or user.employee_role not in KNOWLEDGE_OPERATOR_ROLES: + raise ForbiddenError("仅运营人员可以管理知识库") + return user diff --git a/api/router.py b/api/router.py index 3d132d6..42dbb97 100644 --- a/api/router.py +++ b/api/router.py @@ -3,9 +3,14 @@ """ from fastapi import APIRouter +from api.routers import auth +from api.routers import customer_agent +from api.routers import knowledge from api.routers import auth, product, questionnaire api_router = APIRouter() api_router.include_router(auth.router, prefix="/api", tags=["认证"]) +api_router.include_router(customer_agent.router, prefix="/api/agent/customer", tags=["客服Agent"]) +api_router.include_router(knowledge.router, prefix="/api/knowledge", tags=["知识库"]) api_router.include_router(product.router, prefix="/api", tags=["产品"]) api_router.include_router(questionnaire.router, prefix="/api", tags=["问卷"]) diff --git a/config/database/milvus.py b/config/database/milvus.py index 8004731..67329e6 100644 --- a/config/database/milvus.py +++ b/config/database/milvus.py @@ -24,6 +24,16 @@ def client() -> AsyncMilvusClient: return _client +async def ensure_database(milvus_client: AsyncMilvusClient | None = None) -> None: + """Create the configured project database only when it does not exist.""" + target = milvus_client or client() + database_name = settings.milvus.db_name + if not database_name: + raise RuntimeError("MILVUS_DB must be configured") + if database_name not in await target.list_databases(): + await target.create_database(database_name) + + async def init_db() -> None: client() # 急切建连;失败由注册表重试(config/database/__init__.py) @@ -36,4 +46,4 @@ async def dispose() -> None: async def check_health() -> None: - await client().get_server_version() \ No newline at end of file + await client().get_server_version() diff --git a/config/settings.py b/config/settings.py index 171da69..4349116 100644 --- a/config/settings.py +++ b/config/settings.py @@ -71,12 +71,17 @@ class MilvusCfg(BaseSettings): token: Optional[str] = None # 可缺省(无需鉴权时留空) user: Optional[str] = None password: Optional[str] = None - db_name: Optional[str] = None + db: Optional[str] = None connect_timeout: int # 建连/通道就绪超时(构造是急切连接,必须短) timeout: int # 数据操作超时,调用处可覆盖 model_config = SettingsConfigDict(env_prefix="MILVUS_", env_file=_ENV_FILE, extra="ignore") + @property + def db_name(self) -> Optional[str]: + """Compatibility name used by the Milvus client wrapper.""" + return self.db + class LLMCfg(BaseSettings): """大模型配置:本地 Ollama / OpenAI 兼容 API 双模式见 tool/llm.py。""" @@ -134,4 +139,4 @@ class Settings(BaseSettings): model_config = SettingsConfigDict(env_file=_ENV_FILE, extra="ignore") -settings = Settings() \ No newline at end of file +settings = Settings() diff --git a/main.py b/main.py index b7e6be0..da752ab 100644 --- a/main.py +++ b/main.py @@ -5,6 +5,10 @@ from fastapi import FastAPI from api.router import api_router from config import database +from service.customer_agent.bootstrap import ( + build_default_knowledge_upload_service, + build_default_runtime, +) from utils.exceptions import register_exception_handlers from utils.logger import setup_logging from utils.request_id import RequestIdMiddleware @@ -13,6 +17,8 @@ from utils.request_id import RequestIdMiddleware @asynccontextmanager async def lifespan(app: FastAPI): setup_logging() # 幂等:分级日志 + trace_id + 脱敏 + app.state.customer_agent_runtime = build_default_runtime() + app.state.knowledge_upload_service = build_default_knowledge_upload_service() # 四库会话懒创建:启动不连接任何库,首次访问才建,可手动预热: # asyncio.run(database.init_db()) yield @@ -28,4 +34,8 @@ app.include_router(api_router) @app.get("/") async def root(): - return {"message": "智能公募基金系统 API", "docs": "/docs"} \ No newline at end of file + return {"message": "智能公募基金系统 API", "docs": "/docs"} + +if __name__ == '__main__': + import uvicorn + uvicorn.run('main:app', host="127.0.0.1", port=8000) diff --git a/nl2sql/__init__.py b/nl2sql/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/requirements.txt b/requirements.txt index 0fc9de2..1b34b21 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,3 +1,18 @@ +fastapi +python-multipart +uvicorn +httpx +pyjwt +pydantic +pydantic-settings +sqlalchemy +aiomysql +redis +pymilvus +neo4j +python-dotenv +pypdf + fastapi~=0.141.1 sqlalchemy~=2.0.52 httpx~=0.28.1 diff --git a/sql/schema.sql b/sql/schema.sql index ac56327..a71c2a3 100644 --- a/sql/schema.sql +++ b/sql/schema.sql @@ -265,6 +265,7 @@ CREATE TABLE IF NOT EXISTS conversation_archive ( role VARCHAR(16) NOT NULL COMMENT 'user/assistant/system', content MEDIUMTEXT NULL COMMENT '对话内容', tool_calls JSON NULL COMMENT '工具调用记录 [{"tool":"nl2sql",...}]', + trace_id VARCHAR(64) NULL COMMENT '请求链路追踪ID', create_time DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, KEY idx_session (session_id), KEY idx_user_time (user_id, create_time), @@ -336,7 +337,7 @@ CREATE TABLE IF NOT EXISTS audit_log ( target VARCHAR(128) NULL COMMENT '操作对象(单号/ID)', detail TEXT NULL COMMENT '详情(JSON 字符串)', ip VARCHAR(64) NULL, - trace_id VARCHAR(32) NULL COMMENT '链路号', + trace_id VARCHAR(64) NULL COMMENT '链路号', status VARCHAR(8) NOT NULL DEFAULT '成功' COMMENT '成功/失败', create_time DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, KEY idx_user (user_id), @@ -436,4 +437,4 @@ CREATE TABLE IF NOT EXISTS portfolio_benchmark ( create_time DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP, update_time DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP, UNIQUE KEY uk_risk_level (risk_level) -) COMMENT='组合基准配置表(投顾Agent 再平衡参照,运营可调整)'; \ No newline at end of file +) COMMENT='组合基准配置表(投顾Agent 再平衡参照,运营可调整)'; diff --git a/utils/response.py b/utils/response.py index f61ef0b..53a8062 100644 --- a/utils/response.py +++ b/utils/response.py @@ -15,7 +15,7 @@ class Code: SERVER_ERROR = 500 LLM_FAIL = 1001 # LLM 调用失败 KB_NO_RESULT = 1002 # 知识库检索无结果 - SQL_GEN_FAIL = 1003 # NL2SQL 生成失败 + SQL_GEN_FAIL = 1003 # nl2sql 生成失败 RISK_TRIGGERED = 1004 # 风控规则触发 / 交易被拦截 NOT_SUITABLE = 1005 # 适当性不匹配 -- 2.54.0 From 2004b8fcf4900c4d043262e7973bc35bf7c01229 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E5=8F=B2=E8=8D=92=E4=B8=98?= Date: Fri, 11 Sep 2026 11:02:15 +0800 Subject: [PATCH 2/3] =?UTF-8?q?chore:=20update=20gitignore;=20feat:=20?= =?UTF-8?q?=E6=96=B0=E5=A2=9Ecustomer=5Fagent=E4=B8=9A=E5=8A=A1=E6=A8=A1?= =?UTF-8?q?=E5=9D=97=E4=B8=8Eapi=E8=B7=AF=E7=94=B1?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 68 +++++- .idea/Mutual_Fund.iml | 10 - .idea/misc.xml | 7 - agent/customer_agent/__init__.py | 5 + agent/customer_agent/context.py | 51 +++++ agent/customer_agent/session.py | 67 ++++++ api/routers/customer_agent.py | 82 +++++++ api/routers/knowledge.py | 108 +++++++++ data/mock_knowledge/faq.md | 119 ++++++++++ data/mock_knowledge/fund_product.md | 9 + data/mock_knowledge/policy.md | 9 + rag/chunk_config.py | 30 +++ rag/chunking.py | 114 ++++++++++ rag/cleaning.py | 50 +++++ rag/document_parser.py | 51 +++++ rag/embedding.py | 30 +++ rag/events.py | 19 ++ rag/generation.py | 42 ++++ rag/ingestion.py | 82 +++++++ rag/intent.py | 39 ++++ rag/milvus_collections.py | 46 ++++ rag/milvus_delete.py | 19 ++ rag/mock_ingest.py | 59 +++++ rag/preview.py | 38 ++++ rag/retrieve.py | 158 ++++++++++++++ service/customer_agent/__init__.py | 5 + service/customer_agent/audit.py | 31 +++ service/customer_agent/bootstrap.py | 62 ++++++ service/customer_agent/chat.py | 120 ++++++++++ service/customer_agent/config.py | 23 ++ service/customer_agent/runtime.py | 62 ++++++ service/knowledge_base/__init__.py | 1 + service/knowledge_base/upload.py | 262 ++++++++++++++++++++++ tests/test_agent_context.py | 56 +++++ tests/test_agent_session.py | 88 ++++++++ tests/test_anonymous_agent_service.py | 94 ++++++++ tests/test_chunk_config.py | 29 +++ tests/test_chunking.py | 88 ++++++++ tests/test_cleaning.py | 30 +++ tests/test_customer_agent_audit.py | 24 ++ tests/test_customer_agent_bootstrap.py | 42 ++++ tests/test_customer_agent_config.py | 30 +++ tests/test_customer_agent_packages.py | 14 ++ tests/test_customer_agent_router.py | 99 +++++++++ tests/test_customer_agent_runtime.py | 25 +++ tests/test_document_parser.py | 72 ++++++ tests/test_embedding.py | 35 +++ tests/test_events.py | 22 ++ tests/test_generation.py | 41 ++++ tests/test_ingestion.py | 84 +++++++ tests/test_intent.py | 34 +++ tests/test_knowledge_auth.py | 28 +++ tests/test_knowledge_router.py | 148 +++++++++++++ tests/test_knowledge_upload.py | 290 +++++++++++++++++++++++++ tests/test_milvus_collections.py | 52 +++++ tests/test_milvus_config.py | 34 +++ tests/test_mock_ingest.py | 52 +++++ tests/test_preview.py | 54 +++++ tests/test_retrieve.py | 102 +++++++++ 59 files changed, 3527 insertions(+), 18 deletions(-) delete mode 100644 .idea/Mutual_Fund.iml delete mode 100644 .idea/misc.xml create mode 100644 agent/customer_agent/__init__.py create mode 100644 agent/customer_agent/context.py create mode 100644 agent/customer_agent/session.py create mode 100644 api/routers/customer_agent.py create mode 100644 api/routers/knowledge.py create mode 100644 data/mock_knowledge/faq.md create mode 100644 data/mock_knowledge/fund_product.md create mode 100644 data/mock_knowledge/policy.md create mode 100644 rag/chunk_config.py create mode 100644 rag/chunking.py create mode 100644 rag/cleaning.py create mode 100644 rag/document_parser.py create mode 100644 rag/embedding.py create mode 100644 rag/events.py create mode 100644 rag/generation.py create mode 100644 rag/ingestion.py create mode 100644 rag/intent.py create mode 100644 rag/milvus_collections.py create mode 100644 rag/milvus_delete.py create mode 100644 rag/mock_ingest.py create mode 100644 rag/preview.py create mode 100644 rag/retrieve.py create mode 100644 service/customer_agent/__init__.py create mode 100644 service/customer_agent/audit.py create mode 100644 service/customer_agent/bootstrap.py create mode 100644 service/customer_agent/chat.py create mode 100644 service/customer_agent/config.py create mode 100644 service/customer_agent/runtime.py create mode 100644 service/knowledge_base/__init__.py create mode 100644 service/knowledge_base/upload.py create mode 100644 tests/test_agent_context.py create mode 100644 tests/test_agent_session.py create mode 100644 tests/test_anonymous_agent_service.py create mode 100644 tests/test_chunk_config.py create mode 100644 tests/test_chunking.py create mode 100644 tests/test_cleaning.py create mode 100644 tests/test_customer_agent_audit.py create mode 100644 tests/test_customer_agent_bootstrap.py create mode 100644 tests/test_customer_agent_config.py create mode 100644 tests/test_customer_agent_packages.py create mode 100644 tests/test_customer_agent_router.py create mode 100644 tests/test_customer_agent_runtime.py create mode 100644 tests/test_document_parser.py create mode 100644 tests/test_embedding.py create mode 100644 tests/test_events.py create mode 100644 tests/test_generation.py create mode 100644 tests/test_ingestion.py create mode 100644 tests/test_intent.py create mode 100644 tests/test_knowledge_auth.py create mode 100644 tests/test_knowledge_router.py create mode 100644 tests/test_knowledge_upload.py create mode 100644 tests/test_milvus_collections.py create mode 100644 tests/test_milvus_config.py create mode 100644 tests/test_mock_ingest.py create mode 100644 tests/test_preview.py create mode 100644 tests/test_retrieve.py diff --git a/.gitignore b/.gitignore index 2eea525..4355a67 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,67 @@ -.env \ No newline at end of file +.env +# Python 缓存与编译文件 +__pycache__/ +*.py[cod] +*$py.class +*.so +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +*.egg-info/ +.installed.cfg +*.egg + +# 虚拟环境(非常重要,防止把本地依赖包提交上去) +venv/ +env/ +.idea +.venv/ +ENV/ + +# 日志文件 +*.log + +# 操作系统文件 +.DS_Store +Thumbs.db + +# pytest 测试产物 +.pytest_cache/ +.coverage +coverage.xml +htmlcov/ +junit.xml + +# mypy 类型检查缓存 +.mypy_cache/ +# 向量库、本地数据库、模型文件(大文件不要提交git) +*.db +*.sqlite +*.milvus +*.neo4j +*.bin +*.pth +*.ckpt +*.onnx + +# 本地导出数据、临时csv/excel +data/tmp/ +data/output/ +*.csv +*.xlsx +*.parquet + +# 本地临时需求&开发计划文档 +DK2_客服Agent模块完整开发计划(v1.1).md +DK4_客服Agent需求文档(修复完善版v1.4).md +客服Agent模块 · TODO List(最终修复版v1.2).md diff --git a/.idea/Mutual_Fund.iml b/.idea/Mutual_Fund.iml deleted file mode 100644 index 7b59439..0000000 --- a/.idea/Mutual_Fund.iml +++ /dev/null @@ -1,10 +0,0 @@ - - - - - - - - - - \ No newline at end of file diff --git a/.idea/misc.xml b/.idea/misc.xml deleted file mode 100644 index 34dc9d9..0000000 --- a/.idea/misc.xml +++ /dev/null @@ -1,7 +0,0 @@ - - - - - - \ No newline at end of file diff --git a/agent/customer_agent/__init__.py b/agent/customer_agent/__init__.py new file mode 100644 index 0000000..70fb398 --- /dev/null +++ b/agent/customer_agent/__init__.py @@ -0,0 +1,5 @@ +"""客服 Agent 领域层:会话编排、意图处理和响应策略。""" + +package_name = "customer_agent" + +__all__ = ["package_name"] diff --git a/agent/customer_agent/context.py b/agent/customer_agent/context.py new file mode 100644 index 0000000..020684b --- /dev/null +++ b/agent/customer_agent/context.py @@ -0,0 +1,51 @@ +"""Redis short-term conversation context for the customer-service Agent.""" +from __future__ import annotations + +import json +from inspect import isawaitable + + +async def _config(config_getter, key: str, default): + value = config_getter(key, str(default)) + if isawaitable(value): + value = await value + return type(default)(value) + + +class RedisConversationContext: + def __init__(self, redis, *, config_getter, token_counter=None): + self.redis = redis + self.config_getter = config_getter + self.token_counter = token_counter or (lambda text: max(1, len(text) // 4)) + + async def append(self, session_id: str, role: str, content: str) -> None: + key = f"session:{session_id}:messages" + payload = json.dumps( + {"role": role, "content": content}, ensure_ascii=False + ) + await self.redis.rpush(key, payload) + ttl = await _config(self.config_getter, "agent.customer.session.ttl", 1800) + await self.redis.expire(key, ttl) + await self._trim(key) + + async def get(self, session_id: str) -> list[dict]: + raw_messages = await self.redis.lrange( + f"session:{session_id}:messages", 0, -1 + ) + return [json.loads(raw) for raw in raw_messages if raw] + + async def _trim(self, key: str) -> None: + limit = await _config( + self.config_getter, "agent.customer.session_max_token", 4096 + ) + raw_messages = [raw for raw in await self.redis.lrange(key, 0, -1) if raw] + total = 0 + keep_from = len(raw_messages) + for index in range(len(raw_messages) - 1, -1, -1): + message = json.loads(raw_messages[index]) + total += self.token_counter(message["content"]) + if total > limit: + keep_from = index + 1 + break + keep_from = index + await self.redis.ltrim(key, keep_from, -1) diff --git a/agent/customer_agent/session.py b/agent/customer_agent/session.py new file mode 100644 index 0000000..be92fb9 --- /dev/null +++ b/agent/customer_agent/session.py @@ -0,0 +1,67 @@ +"""Redis-backed anonymous customer-service session primitives.""" +from __future__ import annotations + +import time +import uuid +from inspect import isawaitable + + +class SessionOwnershipError(Exception): + def __init__(self, code: int, message: str): + self.code = code + self.message = message + super().__init__(message) + + +async def _config(config_getter, key: str, default): + value = config_getter(key, str(default)) + if isawaitable(value): + value = await value + return type(default)(value) + + +class AnonymousSessionService: + def __init__(self, redis, *, config_getter, clock=time.time): + self.redis = redis + self.config_getter = config_getter + self.clock = clock + + async def create_session(self) -> str: + session_id = uuid.uuid4().hex + ttl = await _config(self.config_getter, "agent.customer.session.ttl", 1800) + await self.redis.set(f"session:{session_id}", "anonymous", ex=ttl) + # Redis lists cannot be created empty; the marker is ignored by readers. + await self.redis.rpush(f"session:{session_id}:messages", "") + await self.redis.expire(f"session:{session_id}:messages", ttl) + return session_id + + async def verify_session_ownership( + self, + session_id: str, + *, + customer_id: str | None = None, + session_customer_id: str | None = None, + ) -> None: + if not await self.redis.exists(f"session:{session_id}"): + raise SessionOwnershipError(404, "会话不存在或已过期") + if customer_id is not None and customer_id != session_customer_id: + raise SessionOwnershipError(403, "无权访问该会话") + + async def consume_chat_quota(self, session_id: str) -> int | None: + window = await _config( + self.config_getter, "agent.customer.rate_limit.window_sec", 60 + ) + maximum = await _config( + self.config_getter, "agent.customer.rate_limit.max_requests", 20 + ) + key = f"rate:limit:anon:{session_id}:chat" + now = self.clock() + await self.redis.zremrangebyscore(key, 0, (now - window) * 1000) + count = await self.redis.zcard(key) + if count >= maximum: + entries = self.redis.sorted_sets.get(key, []) + oldest = min((score for score, _ in entries), default=now * 1000) + return max(1, int((oldest / 1000 + window) - now)) + await self.redis.zadd(key, {uuid.uuid4().hex: now * 1000}) + await self.redis.expire(key, window) + return None diff --git a/api/routers/customer_agent.py b/api/routers/customer_agent.py new file mode 100644 index 0000000..9b0f47d --- /dev/null +++ b/api/routers/customer_agent.py @@ -0,0 +1,82 @@ +"""Anonymous customer-service Agent HTTP endpoints.""" +from __future__ import annotations + +import json + +from fastapi import APIRouter, Request +from fastapi.responses import JSONResponse, StreamingResponse + +from service.customer_agent.chat import QueryTooLongError +from agent.customer_agent.session import SessionOwnershipError +from utils.request_id import get_request_id, new_request_id +from utils.response import fail, success + + +router = APIRouter() + + +def _runtime(request: Request): + runtime = getattr(request.app.state, "customer_agent_runtime", None) + if runtime is None: + raise RuntimeError("customer agent runtime is not configured") + return runtime + + +@router.post("/session/create") +async def create_session(request: Request): + runtime = _runtime(request) + session_id = await runtime.session_service.create_session() + return success({"session_id": session_id, "customer_id": None}) + + +@router.get("/chat") +async def chat(request: Request, session_id: str, query: str): + runtime = _runtime(request) + trace_id = request.headers.get("X-Trace-Id") or get_request_id() or new_request_id() + try: + await runtime.session_service.verify_session_ownership(session_id) + retry_after = await runtime.session_service.consume_chat_quota(session_id) + if retry_after is not None: + response = fail(429, "请求过于频繁", {"retry_after": retry_after}) + return JSONResponse( + status_code=429, + headers={"Retry-After": str(retry_after)}, + content=response.model_dump(), + ) + result = await runtime.agent.handle(session_id, query, trace_id=trace_id) + except SessionOwnershipError as exc: + return JSONResponse( + status_code=exc.code, + content=fail(exc.code, exc.message).model_dump(), + ) + except QueryTooLongError as exc: + return JSONResponse( + status_code=400, + content=fail(400, str(exc)).model_dump(), + ) + + async def events(): + yield f"data: {json.dumps(result, ensure_ascii=False)}\n\n" + + return StreamingResponse( + events(), media_type="text/event-stream", headers={"X-Trace-Id": trace_id} + ) + + +@router.post("/session/end") +async def end_session(request: Request, body: dict): + runtime = _runtime(request) + session_id = body.get("session_id", "") + try: + await runtime.session_service.verify_session_ownership(session_id) + except SessionOwnershipError as exc: + return JSONResponse( + status_code=exc.code, + content=fail(exc.code, exc.message).model_dump(), + ) + if hasattr(runtime.redis, "delete"): + await runtime.redis.delete( + f"session:{session_id}", f"session:{session_id}:messages", + f"rate:limit:anon:{session_id}:chat", + ) + return success({"session_id": session_id, "archived": False}) diff --git a/api/routers/knowledge.py b/api/routers/knowledge.py new file mode 100644 index 0000000..0ac8477 --- /dev/null +++ b/api/routers/knowledge.py @@ -0,0 +1,108 @@ +"""Knowledge-base upload, preview, and confirmation endpoints.""" +from __future__ import annotations + +from fastapi import APIRouter, Depends, Request +from fastapi.responses import JSONResponse + +from service.knowledge_base.upload import KnowledgeUploadService, UploadValidationError +from api.deps import require_knowledge_operator +from model.sys_user import SysUser +from utils.response import fail, success + + +router = APIRouter() + + +def _service(request: Request) -> KnowledgeUploadService: + service = getattr(request.app.state, "knowledge_upload_service", None) + if service is None: + raise RuntimeError("knowledge upload service is not configured") + return service + + +def _optional_int(value): + return None if value in (None, "") else int(value) + + +@router.get("/documents") +async def list_documents( + request: Request, + _: SysUser = Depends(require_knowledge_operator), +): + return success(await _service(request).list_documents()) + + +@router.get("/documents/{doc_id}") +async def get_document( + request: Request, + doc_id: str, + _: SysUser = Depends(require_knowledge_operator), +): + try: + result = await _service(request).get_document(doc_id) + if result is None: + response = fail(404, "文档不存在") + return JSONResponse(status_code=404, content=response.model_dump()) + return success(result) + except (UploadValidationError, ValueError) as exc: + response = fail(400, str(exc)) + return JSONResponse(status_code=400, content=response.model_dump()) + + +@router.post("/documents/preview") +async def preview_document_upload( + request: Request, + _: SysUser = Depends(require_knowledge_operator), +): + try: + form = await request.form() + upload = form.get("file") + if upload is None or not hasattr(upload, "read"): + raise UploadValidationError("缺少上传文件") + result = await _service(request).preview( + upload.filename, + await upload.read(), + strategy=form.get("strategy", ""), + chunk_size=_optional_int(form.get("chunk_size")), + chunk_overlap=_optional_int(form.get("chunk_overlap")), + ) + return success(result) + except (UploadValidationError, ValueError) as exc: + response = fail(400, str(exc)) + return JSONResponse(status_code=400, content=response.model_dump()) + + +@router.post("/documents/confirm") +async def confirm_document_upload( + request: Request, + _: SysUser = Depends(require_knowledge_operator), +): + try: + body = await request.json() + result = await _service(request).confirm( + upload_id=body.get("upload_id", ""), + title=body.get("title", ""), + doc_id=body.get("doc_id", ""), + collection_name=body.get("collection_name", ""), + strategy=body.get("strategy", ""), + chunk_size=body.get("chunk_size"), + chunk_overlap=body.get("chunk_overlap"), + ) + return success(result) + except (UploadValidationError, ValueError) as exc: + response = fail(400, str(exc)) + return JSONResponse(status_code=400, content=response.model_dump()) + + +@router.delete("/documents/{doc_id}") +async def delete_document( + request: Request, + doc_id: str, + _: SysUser = Depends(require_knowledge_operator), +): + try: + result = await _service(request).delete_document(doc_id) + return success(result) + except (UploadValidationError, ValueError) as exc: + response = fail(400, str(exc)) + return JSONResponse(status_code=400, content=response.model_dump()) diff --git a/data/mock_knowledge/faq.md b/data/mock_knowledge/faq.md new file mode 100644 index 0000000..2c0a7a7 --- /dev/null +++ b/data/mock_knowledge/faq.md @@ -0,0 +1,119 @@ +Q: 什么是开放式基金? +A: 开放式基金的基金份额总数不固定,投资者可以按照基金合同约定申购或赎回。 + +Q: 基金交易的确认时间是什么时候? +A: 通常以交易日收市时间为界,具体确认规则以基金合同和销售机构公告为准。 + +Q: 什么是货币基金? +A: 货币基金主要投资于短期货币市场工具,具体投资范围以基金合同为准。 + +Q: 什么是债券基金? +A: 债券基金主要投资于债券等固定收益类资产,仍可能受到市场波动影响。 + +Q: 什么是股票基金? +A: 股票基金主要投资于股票,净值波动通常与股票市场变化相关。 + +Q: 基金净值每天都会更新吗? +A: 开放式基金通常在交易日公布净值,具体时间以基金管理人公告为准。 + +Q: 基金申购是什么意思? +A: 申购是投资者按照基金份额净值购买开放式基金份额的行为。 + +Q: 基金赎回是什么意思? +A: 赎回是投资者向基金管理人申请卖出基金份额并取得资金的行为。 + +Q: 基金转换是什么? +A: 基金转换是将持有的一只基金份额转换为同一基金管理人旗下另一只基金份额。 + +Q: 基金定投是什么? +A: 基金定投是按约定周期和金额持续投资指定基金的方式。 + +Q: 基金分红会增加收益吗? +A: 基金分红通常是基金资产的一部分以现金或再投资形式分配给投资者,不等同于额外收益。 + +Q: 基金分红有哪些方式? +A: 常见方式包括现金分红和红利再投资,具体以基金合同及投资者设置为准。 + +Q: 基金费用包括哪些? +A: 基金费用可能包括管理费、托管费、销售服务费以及申购赎回费用等。 + +Q: 申购费什么时候收取? +A: 申购费通常在申购基金时收取,具体费率和收取方式以基金公告为准。 + +Q: 赎回费和持有时间有关吗? +A: 部分基金的赎回费率会根据持有期限设置差异,具体以基金合同为准。 + +Q: 基金有最低申购金额吗? +A: 不同基金和销售渠道可能设置不同的最低申购金额,应以对应页面和公告为准。 + +Q: 基金可以随时买卖吗? +A: 开放式基金通常可在开放日办理业务,封闭式基金和特殊产品规则可能不同。 + +Q: 节假日可以申购基金吗? +A: 节假日期间是否受理以及确认时间取决于基金和销售机构的业务安排。 + +Q: 基金交易日怎么判断? +A: 基金交易日通常参考证券交易所交易日安排,具体以基金管理人公告为准。 + +Q: 基金净值估算准确吗? +A: 净值估算仅为参考信息,最终净值以基金管理人公布的数据为准。 + +Q: 基金投资有本金保障吗? +A: 基金投资通常不承诺保本,投资者应阅读基金合同并自行承担相应风险。 + +Q: 如何查看基金风险等级? +A: 可在基金详情页、基金合同或销售机构风险揭示文件中查看风险等级。 + +Q: 风险承受能力评估是什么? +A: 风险承受能力评估用于了解投资者的风险偏好和承受能力,结果不代表收益承诺。 + +Q: 基金适合短期投资吗? +A: 是否适合短期投资取决于产品特征、投资目标和个人风险承受能力。 + +Q: 基金可以撤销交易吗? +A: 部分交易在规定时间内可能支持撤单,能否撤销应以销售机构规则为准。 + +Q: 基金认购和申购有什么区别? +A: 认购通常指募集期购买新基金,申购通常指基金成立后购买开放式基金份额。 + +Q: 新基金募集期多长? +A: 募集期长短因产品而异,以基金管理人发布的募集公告为准。 + +Q: 基金成立后多久可以赎回? +A: 具体开放赎回时间取决于基金合同和公告,封闭运作产品可能有特别安排。 + +Q: 基金规模会变化吗? +A: 开放式基金规模会随着申购、赎回和净值变化而变化。 + +Q: 基金经理会影响基金表现吗? +A: 基金经理是影响投资管理的重要因素之一,但不代表未来业绩保证。 + +Q: 基金历史业绩能代表未来吗? +A: 历史业绩仅供参考,不能代表未来收益,也不构成投资建议。 + +Q: 基金公告在哪里查看? +A: 可通过基金管理人官网、销售机构页面或依法披露的信息渠道查看。 + +Q: 基金合同是什么? +A: 基金合同约定基金运作方式、投资范围、费用和各方权利义务等重要事项。 + +Q: 基金招募说明书有什么作用? +A: 招募说明书介绍基金产品的重要信息,投资者应在投资前认真阅读。 + +Q: 什么是基金托管人? +A: 基金托管人依法保管基金财产并履行监督等职责,具体职责以法律法规为准。 + +Q: 什么是基金管理人? +A: 基金管理人负责基金募集、投资管理和信息披露等工作,并依法承担相应责任。 + +Q: 基金账户和银行卡一样吗? +A: 基金账户用于记录基金交易和份额信息,不等同于银行结算账户。 + +Q: 如何查询我的基金份额? +A: 可通过销售机构的账户或交易页面查询,具体查询方式以渠道功能为准。 + +Q: 基金赎回后资金多久到账? +A: 到账时间因基金类型、交易渠道和节假日安排而异,应以页面提示和公告为准。 + +Q: 基金暂停赎回是什么意思? +A: 暂停赎回表示在特定情形下暂时停止受理赎回申请,具体以相关公告为准。 diff --git a/data/mock_knowledge/fund_product.md b/data/mock_knowledge/fund_product.md new file mode 100644 index 0000000..31eafee --- /dev/null +++ b/data/mock_knowledge/fund_product.md @@ -0,0 +1,9 @@ +# 示例基金产品 + +## 产品概况 + +本基金为示例产品,主要用于客服 Agent 的知识库联调和检索验证。 + +## 风险收益特征 + +基金投资存在风险,示例内容不构成任何投资建议,实际信息以正式法律文件为准。 diff --git a/data/mock_knowledge/policy.md b/data/mock_knowledge/policy.md new file mode 100644 index 0000000..b5767ec --- /dev/null +++ b/data/mock_knowledge/policy.md @@ -0,0 +1,9 @@ +# 示例政策法规 + +## 信息披露 + +基金管理人应按照法律法规和监管要求履行信息披露义务。 + +## 投资者权益 + +投资者依法享有查询、赎回以及获取相关信息的权利。 diff --git a/rag/chunk_config.py b/rag/chunk_config.py new file mode 100644 index 0000000..881a22b --- /dev/null +++ b/rag/chunk_config.py @@ -0,0 +1,30 @@ +"""Project-level chunking defaults and validation.""" +from __future__ import annotations + +from dataclasses import dataclass + + +DEFAULT_CHUNK_SIZE = 512 +DEFAULT_CHUNK_OVERLAP = 64 + + +@dataclass(frozen=True) +class ChunkConfig: + size: int + overlap: int + + +def resolve_chunk_config( + *, + chunk_size: int | None = None, + chunk_overlap: int | None = None, +) -> ChunkConfig: + size = DEFAULT_CHUNK_SIZE if chunk_size is None else chunk_size + overlap = DEFAULT_CHUNK_OVERLAP if chunk_overlap is None else chunk_overlap + if size <= 0: + raise ValueError("chunk_size must be greater than zero") + if overlap < 0: + raise ValueError("chunk_overlap cannot be negative") + if overlap >= size: + raise ValueError("chunk_overlap must be smaller than chunk_size") + return ChunkConfig(size=size, overlap=overlap) diff --git a/rag/chunking.py b/rag/chunking.py new file mode 100644 index 0000000..3d670f0 --- /dev/null +++ b/rag/chunking.py @@ -0,0 +1,114 @@ +"""Document chunking strategies used before Milvus ingestion.""" +from __future__ import annotations + +import re +from dataclasses import dataclass + +from rag.chunk_config import ChunkConfig, resolve_chunk_config + + +SUPPORTED_STRATEGIES = {"default", "qa_pair", "chapter_semantic"} + + +@dataclass(frozen=True) +class Chunk: + text: str + section_title: str | None = None + + +@dataclass(frozen=True) +class ChunkingResult: + chunks: list[Chunk] + requested_strategy: str + actual_strategy: str + degraded: bool = False + warning: str | None = None + + +def chunk_document( + text: str, + strategy: str, + *, + config: ChunkConfig | None = None, +) -> ChunkingResult: + if strategy not in SUPPORTED_STRATEGIES: + raise ValueError(f"Unsupported chunk strategy: {strategy}") + config = config or resolve_chunk_config() + if strategy == "default": + return ChunkingResult(_default_chunks(text, config), strategy, strategy) + if strategy == "qa_pair": + return ChunkingResult(_qa_pair_chunks(text, config), strategy, strategy) + return _chapter_chunks(text, config) + + +def _default_chunks(text: str, config: ChunkConfig) -> list[Chunk]: + paragraphs = [part.strip() for part in re.split(r"\n\s*\n", text) if part.strip()] + chunks: list[Chunk] = [] + for paragraph in paragraphs: + chunks.extend(Chunk(part) for part in _split_long_text(paragraph, config)) + return chunks + + +def _split_long_text(text: str, config: ChunkConfig) -> list[str]: + if len(text) <= config.size: + return [text] + chunks = [] + start = 0 + while start < len(text): + end = min(start + config.size, len(text)) + chunks.append(text[start:end]) + if end == len(text): + break + start = end - config.overlap + return chunks + + +def _qa_pair_chunks(text: str, config: ChunkConfig) -> list[Chunk]: + matches = list( + re.finditer( + r"^\s*Q:\s*(.*?)\r?\n\s*A:\s*(.*?)(?=^\s*Q:\s*|\Z)", + text, + flags=re.MULTILINE | re.DOTALL, + ) + ) + if not matches or any(match.group(0).strip() == "" for match in matches): + raise ValueError("未识别FAQ问答格式,请确认文档包含Q:/A:标记") + if any(not match.group(1).strip() or not match.group(2).strip() for match in matches): + raise ValueError("FAQ问答对必须同时包含问题和答案") + chunks = [ + Chunk(f"question: {match.group(1).strip()}\nanswer: {match.group(2).strip()}") + for match in matches + ] + if any(len(chunk.text) > config.size for chunk in chunks): + raise ValueError("FAQ问答对超过chunk_size,整份文档拒绝上传") + return chunks + + +def _chapter_chunks(text: str, config: ChunkConfig) -> ChunkingResult: + heading_pattern = re.compile(r"^(#{1,3})\s+(.+?)\s*$", re.MULTILINE) + headings = list(heading_pattern.finditer(text)) + if not headings: + return ChunkingResult( + _default_chunks(text, config), + "chapter_semantic", + "default", + degraded=True, + warning="未发现Markdown标题,已降级为default策略", + ) + + chunks: list[Chunk] = [] + title_stack: list[str] = [] + for index, heading in enumerate(headings): + level = len(heading.group(1)) + title = heading.group(2).strip() + title_stack = title_stack[: level - 1] + [title] + body_start = heading.end() + body_end = headings[index + 1].start() if index + 1 < len(headings) else len(text) + body = text[body_start:body_end].strip() + if not body: + continue + section_title = " > ".join(title_stack) + prefix = f"【章节:{section_title}】\n" + for part in _split_long_text(body, config): + chunks.append(Chunk(prefix + part, section_title=section_title)) + return ChunkingResult(chunks, "chapter_semantic", "chapter_semantic") diff --git a/rag/cleaning.py b/rag/cleaning.py new file mode 100644 index 0000000..ec2eb56 --- /dev/null +++ b/rag/cleaning.py @@ -0,0 +1,50 @@ +"""Conservative text normalization before chunking knowledge documents.""" +from __future__ import annotations + +import re +from dataclasses import dataclass + + +_PAGE_MARKER = re.compile(r"^\s*(?:第\s*\d+\s*页|page\s+\d+)\s*$", re.I) + + +@dataclass(frozen=True) +class CleanedDocument: + text: str + warnings: list[str] + changed: bool + + +def clean_document_text(text: str) -> CleanedDocument: + """Normalize parser output without changing business values or punctuation.""" + if not isinstance(text, str): + raise TypeError("document text must be a string") + + original = text + warnings: list[str] = [] + text = text.replace("\ufeff", "").replace("\r\n", "\n").replace("\r", "\n") + if text != original: + warnings.append("已统一编码标记和换行符") + + cleaned_lines: list[str] = [] + removed_page_markers = 0 + for line in text.split("\n"): + line = line.replace("\u00a0", " ") + line = "".join(char for char in line if char in "\t\n" or ord(char) >= 32) + if _PAGE_MARKER.match(line): + removed_page_markers += 1 + continue + line = re.sub(r"[ \t]+", " ", line).strip() + cleaned_lines.append(line) + if removed_page_markers: + warnings.append(f"已移除{removed_page_markers}个独立页码标记") + + normalized = "\n".join(cleaned_lines) + normalized = re.sub(r"\n{3,}", "\n\n", normalized).strip() + if normalized != text.strip(): + warnings.append("已清理多余空白、控制字符或空行") + return CleanedDocument( + text=normalized, + warnings=warnings, + changed=normalized != original, + ) diff --git a/rag/document_parser.py b/rag/document_parser.py new file mode 100644 index 0000000..6407824 --- /dev/null +++ b/rag/document_parser.py @@ -0,0 +1,51 @@ +"""Knowledge document parsing for supported upload formats.""" +from __future__ import annotations + +from pathlib import Path +from xml.etree import ElementTree +from zipfile import BadZipFile, ZipFile + +from pypdf import PdfReader + +SUPPORTED_EXTENSIONS = {".txt", ".md", ".docx", ".pdf"} +_WORD_NS = "{http://schemas.openxmlformats.org/wordprocessingml/2006/main}" + + +def parse_document(path: str | Path) -> str: + """Return document text for a supported knowledge-base upload.""" + document_path = Path(path) + suffix = document_path.suffix.lower() + if suffix not in SUPPORTED_EXTENSIONS: + raise ValueError(f"Unsupported document type: {suffix or ''}") + if not document_path.is_file(): + raise FileNotFoundError(document_path) + + if suffix in {".txt", ".md"}: + return document_path.read_text(encoding="utf-8-sig") + if suffix == ".pdf": + return _parse_pdf(document_path) + return _parse_docx(document_path) + + +def _parse_docx(path: Path) -> str: + try: + with ZipFile(path) as archive: + xml = archive.read("word/document.xml") + root = ElementTree.fromstring(xml) + except (BadZipFile, KeyError, ElementTree.ParseError) as exc: + raise ValueError(f"Invalid DOCX document: {path}") from exc + + paragraphs = [] + for paragraph in root.iter(f"{_WORD_NS}p"): + text = "".join(node.text or "" for node in paragraph.iter(f"{_WORD_NS}t")) + if text: + paragraphs.append(text) + return "\n".join(paragraphs) + + +def _parse_pdf(path: Path) -> str: + try: + reader = PdfReader(str(path)) + return "\n".join(page.extract_text() or "" for page in reader.pages).strip() + except Exception as exc: # pypdf raises format-specific exceptions + raise ValueError(f"Invalid PDF document: {path}") from exc diff --git a/rag/embedding.py b/rag/embedding.py new file mode 100644 index 0000000..eb9076f --- /dev/null +++ b/rag/embedding.py @@ -0,0 +1,30 @@ +"""Embedding provider wrapper for Milvus ingestion and retrieval.""" +from __future__ import annotations + +import logging + +from tool.llm import llm + + +logger = logging.getLogger("rag.embedding") +EMBEDDING_DIMENSION = 768 + + +class EmbeddingError(RuntimeError): + """Raised when the embedding provider fails or returns invalid vectors.""" + + +async def embed_texts(texts: list[str], *, client=None) -> list[list[float]]: + if not texts: + return [] + provider = client or llm + try: + vectors = await provider.embed(texts) + except Exception as exc: # provider-specific exceptions are normalized here + logger.exception("embedding provider failed") + raise EmbeddingError("Embedding service unavailable") from exc + if len(vectors) != len(texts) or any( + len(vector) != EMBEDDING_DIMENSION for vector in vectors + ): + raise EmbeddingError(f"Embedding dimension must be {EMBEDDING_DIMENSION}") + return vectors diff --git a/rag/events.py b/rag/events.py new file mode 100644 index 0000000..534cede --- /dev/null +++ b/rag/events.py @@ -0,0 +1,19 @@ +"""Knowledge-base events published by the operations workflow.""" +from __future__ import annotations + + +KNOWLEDGE_UPDATE_EVENT = "event:knowledge_update" + + +async def publish_knowledge_update( + publisher, *, doc_id: str, collection_name: str, chunk_count: int, action: str | None = None +) -> None: + """Publish only after the caller's atomic Milvus write has succeeded.""" + payload = { + "doc_id": doc_id, + "collection_name": collection_name, + "chunk_count": chunk_count, + } + if action is not None: + payload["action"] = action + await publisher(KNOWLEDGE_UPDATE_EVENT, payload) diff --git a/rag/generation.py b/rag/generation.py new file mode 100644 index 0000000..6ee0dde --- /dev/null +++ b/rag/generation.py @@ -0,0 +1,42 @@ +"""LLM answer generation with the客服 Agent fallback chain.""" +from __future__ import annotations + +import logging +from inspect import isawaitable + + +logger = logging.getLogger("rag.generation") + + +async def _config(config_getter, key: str, default=None): + value = config_getter(key, default) + if isawaitable(value): + value = await value + return value + + +async def generate_answer( + messages: list[dict], + *, + llm_client, + config_getter, + primary_model: str | None = None, +) -> str: + fallback_model = await _config( + config_getter, "agent.customer.llm.fallback_model", "" + ) + template = await _config( + config_getter, "agent.customer.template.system_busy", None + ) + models = [primary_model] if primary_model else [None] + if fallback_model and fallback_model not in models: + models.append(fallback_model) + for model in models: + try: + kwargs = {} if model is None else {"model": model} + return await llm_client.chat(messages, **kwargs) + except Exception: + logger.exception("LLM generation failed for model=%s", model or "default") + if not template: + raise RuntimeError("system busy template is not configured") + return template diff --git a/rag/ingestion.py b/rag/ingestion.py new file mode 100644 index 0000000..02a4516 --- /dev/null +++ b/rag/ingestion.py @@ -0,0 +1,82 @@ +"""Validated, all-or-nothing ingestion of a document into Milvus.""" +from __future__ import annotations + +import logging +from uuid import NAMESPACE_URL, uuid5 + +from rag.chunk_config import ChunkConfig, resolve_chunk_config +from rag.chunking import chunk_document +from rag.cleaning import clean_document_text +from rag.embedding import EMBEDDING_DIMENSION, embed_texts +from rag.milvus_delete import _escape_filter_value + + +logger = logging.getLogger("rag.ingestion") + + +async def ingest_document_atomic( + document_text: str, + doc_id: str, + title: str, + collection_name: str, + strategy: str, + *, + milvus_client, + embedder=embed_texts, + config: ChunkConfig | None = None, +) -> dict: + """Validate, embed, and insert one document without leaving partial rows.""" + if not doc_id or not doc_id.strip(): + raise ValueError("doc_id must not be empty") + if not title or not title.strip(): + raise ValueError("title must not be empty") + if not collection_name or not collection_name.strip(): + raise ValueError("collection_name must not be empty") + + resolved_config = config or resolve_chunk_config() + cleaned = clean_document_text(document_text) + chunking = chunk_document(cleaned.text, strategy, config=resolved_config) + if not chunking.chunks: + raise ValueError("document produced no chunks") + + try: + vectors = await embedder([chunk.text for chunk in chunking.chunks]) + if len(vectors) != len(chunking.chunks) or any( + len(vector) != EMBEDDING_DIMENSION for vector in vectors + ): + raise ValueError(f"Embedding dimension must be {EMBEDDING_DIMENSION}") + rows = [ + { + "chunk_id": str(uuid5(NAMESPACE_URL, f"{doc_id}:{index}")), + "doc_id": doc_id, + "title": title, + "section_title": chunk.section_title or "", + "text": chunk.text, + "strategy": strategy, + "vector": vector, + } + for index, (chunk, vector) in enumerate(zip(chunking.chunks, vectors)) + ] + await milvus_client.insert(collection_name=collection_name, data=rows) + except Exception: + logger.exception("atomic Milvus ingestion failed for doc_id=%s", doc_id) + try: + await milvus_client.delete( + collection_name=collection_name, + filter=f'doc_id == "{_escape_filter_value(doc_id)}"', + ) + except Exception: + logger.exception("failed to clean up document rows for doc_id=%s", doc_id) + raise + + return { + "doc_id": doc_id, + "collection_name": collection_name, + "strategy": strategy, + "actual_strategy": chunking.actual_strategy, + "degraded": chunking.degraded, + "warning": chunking.warning, + "cleaning_changed": cleaned.changed, + "cleaning_warnings": cleaned.warnings, + "chunk_count": len(rows), + } diff --git a/rag/intent.py b/rag/intent.py new file mode 100644 index 0000000..cf51eff --- /dev/null +++ b/rag/intent.py @@ -0,0 +1,39 @@ +"""Customer-service intent recognition contract.""" +from __future__ import annotations + +import logging +from enum import StrEnum + + +logger = logging.getLogger("rag.intent") + + +class Intent(StrEnum): + GUIDE_PURCHASE = "guide_purchase" + WANT_ADVISOR = "want_advisor" + NL2SQL_REQUEST = "nl2sql_request" + COMPLAIN = "complain" + KNOWLEDGE_QA = "knowledge_qa" + NO_MATCH = "no_match" + + +INTENT_VALUES = frozenset(item.value for item in Intent) + + +async def intent_recognize(query: str, *, llm_client) -> Intent: + if not query or not query.strip(): + return Intent.NO_MATCH + messages = [ + { + "role": "system", + "content": "仅返回一个意图枚举值,不要输出解释。", + }, + {"role": "user", "content": query}, + ] + try: + raw = await llm_client.chat(messages) + except Exception: + logger.exception("intent recognition failed") + return Intent.NO_MATCH + value = raw.strip().strip('`').lower() + return Intent(value) if value in INTENT_VALUES else Intent.NO_MATCH diff --git a/rag/milvus_collections.py b/rag/milvus_collections.py new file mode 100644 index 0000000..628c827 --- /dev/null +++ b/rag/milvus_collections.py @@ -0,0 +1,46 @@ +"""Milvus collection schemas for the mutual_fund knowledge base.""" +from __future__ import annotations + +from pymilvus import AsyncMilvusClient, DataType + +from config.database.milvus import client as configured_milvus_client + + +KNOWLEDGE_COLLECTIONS = ("fin_faq", "fin_fund_doc", "fin_policy") +EMBEDDING_DIMENSION = 768 + + +def build_knowledge_schema(): + schema = AsyncMilvusClient.create_schema(auto_id=False, enable_dynamic_field=False) + schema.add_field("chunk_id", DataType.VARCHAR, is_primary=True, max_length=128) + schema.add_field("doc_id", DataType.VARCHAR, max_length=128) + schema.add_field("title", DataType.VARCHAR, max_length=512) + schema.add_field("section_title", DataType.VARCHAR, max_length=1024) + schema.add_field("text", DataType.VARCHAR, max_length=65535) + schema.add_field("strategy", DataType.VARCHAR, max_length=32) + schema.add_field("vector", DataType.FLOAT_VECTOR, dim=EMBEDDING_DIMENSION) + return schema + + +def build_knowledge_index_params(): + params = AsyncMilvusClient.prepare_index_params() + params.add_index( + field_name="vector", + index_type="HNSW", + metric_type="COSINE", + params={"M": 16, "efConstruction": 200}, + ) + return params + + +async def ensure_collections(milvus_client: AsyncMilvusClient | None = None) -> None: + client = milvus_client or configured_milvus_client() + schema = build_knowledge_schema() + index_params = build_knowledge_index_params() + for collection_name in KNOWLEDGE_COLLECTIONS: + if not await client.has_collection(collection_name): + await client.create_collection( + collection_name=collection_name, + schema=schema, + index_params=index_params, + ) diff --git a/rag/milvus_delete.py b/rag/milvus_delete.py new file mode 100644 index 0000000..0fd688b --- /dev/null +++ b/rag/milvus_delete.py @@ -0,0 +1,19 @@ +"""Milvus knowledge-document deletion helpers.""" +from __future__ import annotations + +from rag.milvus_collections import KNOWLEDGE_COLLECTIONS +from config.database.milvus import client as configured_milvus_client + + +def _escape_filter_value(value: str) -> str: + return value.replace("\\", "\\\\").replace('"', '\\"') + + +async def delete_document_vectors(doc_id: str, *, milvus_client=None) -> None: + """Delete every knowledge chunk with ``doc_id`` from project collections.""" + if not doc_id or not doc_id.strip(): + raise ValueError("doc_id must not be empty") + client = milvus_client or configured_milvus_client() + expression = f'doc_id == "{_escape_filter_value(doc_id)}"' + for collection_name in KNOWLEDGE_COLLECTIONS: + await client.delete(collection_name=collection_name, filter=expression) diff --git a/rag/mock_ingest.py b/rag/mock_ingest.py new file mode 100644 index 0000000..3fd05dd --- /dev/null +++ b/rag/mock_ingest.py @@ -0,0 +1,59 @@ +"""Development-only Mock knowledge documents and Milvus importer.""" +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from uuid import uuid5, NAMESPACE_URL + +from rag.chunk_config import resolve_chunk_config +from rag.chunking import chunk_document +from rag.cleaning import clean_document_text +from rag.document_parser import parse_document +from rag.embedding import embed_texts + + +_ROOT = Path(__file__).resolve().parent.parent +MOCK_DATA_DIR = _ROOT / "data" / "mock_knowledge" + + +@dataclass(frozen=True) +class MockDocument: + doc_id: str + title: str + collection: str + strategy: str + path: Path + + +MOCK_DOCUMENTS = ( + MockDocument("mock-faq-001", "常见问题示例", "fin_faq", "qa_pair", MOCK_DATA_DIR / "faq.md"), + MockDocument("mock-fund-doc-001", "基金产品说明示例", "fin_fund_doc", "chapter_semantic", MOCK_DATA_DIR / "fund_product.md"), + MockDocument("mock-policy-001", "政策法规示例", "fin_policy", "chapter_semantic", MOCK_DATA_DIR / "policy.md"), +) + + +async def ingest_mock_documents(milvus_client, *, embedder=embed_texts) -> None: + """Import all checked-in Markdown samples into their Milvus collections.""" + config = resolve_chunk_config() + for document in MOCK_DOCUMENTS: + result = chunk_document( + clean_document_text(parse_document(document.path)).text, + document.strategy, + config=config, + ) + vectors = await embedder([chunk.text for chunk in result.chunks]) + rows = [] + for index, (chunk, vector) in enumerate(zip(result.chunks, vectors)): + chunk_id = str(uuid5(NAMESPACE_URL, f"{document.doc_id}:{index}")) + rows.append( + { + "chunk_id": chunk_id, + "doc_id": document.doc_id, + "title": document.title, + "section_title": chunk.section_title or "", + "text": chunk.text, + "strategy": document.strategy, + "vector": vector, + } + ) + await milvus_client.insert(collection_name=document.collection, data=rows) diff --git a/rag/preview.py b/rag/preview.py new file mode 100644 index 0000000..4e83451 --- /dev/null +++ b/rag/preview.py @@ -0,0 +1,38 @@ +"""Preview service for manually selected knowledge-base chunking.""" +from __future__ import annotations + +from pathlib import Path + +from rag.chunk_config import resolve_chunk_config +from rag.chunking import SUPPORTED_STRATEGIES, chunk_document +from rag.cleaning import clean_document_text +from rag.document_parser import parse_document + + +def preview_document( + path: str | Path, + *, + strategy: str, + chunk_size: int | None = None, + chunk_overlap: int | None = None, +) -> dict: + """Parse and preview a document without writing anything to Milvus.""" + if strategy not in SUPPORTED_STRATEGIES: + raise ValueError(f"Unsupported chunk strategy: {strategy}") + config = resolve_chunk_config(chunk_size=chunk_size, chunk_overlap=chunk_overlap) + cleaned = clean_document_text(parse_document(path)) + result = chunk_document(cleaned.text, strategy, config=config) + return { + "strategy": result.requested_strategy, + "actual_strategy": result.actual_strategy, + "degraded": result.degraded, + "warning": result.warning, + "cleaning_changed": cleaned.changed, + "cleaning_warnings": cleaned.warnings, + "chunk_size": config.size, + "chunk_overlap": config.overlap, + "chunks": [ + {"text": chunk.text, "section_title": chunk.section_title} + for chunk in result.chunks + ], + } diff --git a/rag/retrieve.py b/rag/retrieve.py new file mode 100644 index 0000000..8c7a02d --- /dev/null +++ b/rag/retrieve.py @@ -0,0 +1,158 @@ +"""Milvus-only retrieval primitives for the customer-service RAG layer.""" +from __future__ import annotations + +from inspect import isawaitable +import logging + +from rag.embedding import embed_texts + + +logger = logging.getLogger("rag.retrieve") + + +BUSINESS_RETRIEVAL = ( + ("fin_faq", "faq", 3, 0.75), + ("fin_fund_doc", "funddoc", 5, 0.70), + ("fin_policy", "policy", 5, 0.70), +) + + +async def _config_value(config_getter, key: str, default): + value = config_getter(key, str(default)) + if isawaitable(value): + value = await value + return type(default)(value) + + +async def retrieve_candidates( + query: str, + customer_id: str | None, + *, + milvus_client, + embedder=embed_texts, + config_getter, +) -> list[dict]: + if not query or not query.strip(): + return [] + vectors = await embedder([query]) + vector = vectors[0] + plans = list(BUSINESS_RETRIEVAL) + if customer_id: + plans.append(("customer_memory", "memory", 5, 0.60)) + + candidates = [] + for collection_name, key_suffix, default_topk, default_threshold in plans: + topk = await _config_value( + config_getter, + f"agent.customer.rag.topk.{key_suffix}", + default_topk, + ) + threshold = await _config_value( + config_getter, + f"agent.customer.rag.threshold.{key_suffix}", + default_threshold, + ) + expression = "" + if collection_name == "customer_memory": + escaped = customer_id.replace("\\", "\\\\").replace('"', '\\"') + expression = f'customer_id == "{escaped}"' + result = await milvus_client.search( + collection_name=collection_name, + data=[vector], + limit=topk, + filter=expression, + output_fields=["doc_id", "title", "section_title", "text", "strategy"], + ) + candidates.append( + { + "collection_name": collection_name, + "threshold": threshold, + "results": result, + } + ) + return candidates + + +def _flatten_hits(results): + for batch in results or []: + if isinstance(batch, dict): + yield batch + else: + yield from batch or [] + + +def _format_candidates(candidates) -> list[dict]: + sources = [] + for candidate in candidates: + threshold = candidate["threshold"] + for hit in _flatten_hits(candidate["results"]): + entity = hit.get("entity") or hit + score = hit.get("distance", hit.get("score")) + if score is None or score < threshold: + continue + sources.append( + { + "doc_id": entity.get("doc_id", ""), + "title": entity.get("title", ""), + "section_title": entity.get("section_title") or None, + "chunk_text": entity.get("text", ""), + "score": score, + } + ) + sources.sort(key=lambda source: source["score"], reverse=True) + return sources + + +async def rag_retrieve( + query: str, + customer_id: str | None, + *, + milvus_client, + embedder=embed_texts, + config_getter, +) -> list[dict]: + """Return the stable source contract consumed by客服 Agent.""" + try: + candidates = await retrieve_candidates( + query, + customer_id, + milvus_client=milvus_client, + embedder=embedder, + config_getter=config_getter, + ) + except Exception: + logger.exception("RAG retrieval failed") + return [] + return _format_candidates(candidates) + + +async def retrieve_with_status( + query: str, + customer_id: str | None, + *, + milvus_client, + embedder=embed_texts, + config_getter, +) -> dict: + """Expose operational status while keeping failed source lists empty.""" + try: + vectors = await embedder([query]) + except Exception: + logger.exception("RAG embedding failed") + return {"status": "embedding_failed", "sources": []} + + async def reuse_vector(_texts): + return vectors + + try: + candidates = await retrieve_candidates( + query, + customer_id, + milvus_client=milvus_client, + embedder=reuse_vector, + config_getter=config_getter, + ) + except Exception: + logger.exception("Milvus retrieval failed") + return {"status": "milvus_unavailable", "sources": []} + return {"status": "ok", "sources": _format_candidates(candidates)} diff --git a/service/customer_agent/__init__.py b/service/customer_agent/__init__.py new file mode 100644 index 0000000..7eeed52 --- /dev/null +++ b/service/customer_agent/__init__.py @@ -0,0 +1,5 @@ +"""客服 Agent 服务层:对外提供匿名/登录客服业务服务。""" + +package_name = "customer_agent" + +__all__ = ["package_name"] diff --git a/service/customer_agent/audit.py b/service/customer_agent/audit.py new file mode 100644 index 0000000..d9a95b8 --- /dev/null +++ b/service/customer_agent/audit.py @@ -0,0 +1,31 @@ +"""Audit-log persistence adapter for anonymous customer-service events.""" +from __future__ import annotations + +import json + +from sqlalchemy import text + + +async def write_anonymous_sensitive_audit(db, *, session_id: str, trace_id: str) -> None: + statement = text( + """ + INSERT INTO audit_log + (user_id, username, module, action, target, detail, trace_id, status) + VALUES + (:user_id, :username, :module, :action, :target, :detail, :trace_id, :status) + """ + ) + await db.execute( + statement, + { + "user_id": None, + "username": None, + "module": "customer_agent", + "action": "anon_sensitive_input", + "target": session_id, + "detail": json.dumps({"session_id": session_id}, ensure_ascii=False), + "trace_id": trace_id, + "status": "成功", + }, + ) + await db.commit() diff --git a/service/customer_agent/bootstrap.py b/service/customer_agent/bootstrap.py new file mode 100644 index 0000000..423562a --- /dev/null +++ b/service/customer_agent/bootstrap.py @@ -0,0 +1,62 @@ +"""Default application wiring for the客服 Agent runtime.""" +from __future__ import annotations + +import json +from pathlib import Path + +from config.database.milvus import client as milvus_client +from config.database.mysql import get_session_factory +from config.database.redis import client as redis_client +from service.customer_agent.config import DatabaseConfigProvider +from service.customer_agent.runtime import build_anonymous_runtime +from service.knowledge_base.upload import KnowledgeUploadService +from rag.embedding import embed_texts +from rag.milvus_collections import KNOWLEDGE_COLLECTIONS +from rag.milvus_delete import _escape_filter_value +from tool.llm import llm as llm_client + + +def build_default_runtime(): + provider = DatabaseConfigProvider(session_factory=get_session_factory()) + return build_anonymous_runtime( + redis=redis_client(), + milvus_client=milvus_client(), + llm_client=llm_client, + config_getter=provider.get, + audit_writer=provider.write_audit, + ) + + +async def document_exists_in_milvus(milvus, doc_id: str) -> bool: + """Check duplicate document IDs across all project knowledge collections.""" + escaped_doc_id = _escape_filter_value(doc_id) + for collection_name in KNOWLEDGE_COLLECTIONS: + rows = await milvus.query( + collection_name=collection_name, + filter=f'doc_id == "{escaped_doc_id}"', + output_fields=["doc_id"], + limit=1, + ) + if rows: + return True + return False + + +def build_default_knowledge_upload_service(): + provider = DatabaseConfigProvider(session_factory=get_session_factory()) + redis = redis_client() + milvus = milvus_client() + + async def publish_event(event_name, payload): + await redis.publish( + event_name, + json.dumps(payload, ensure_ascii=False), + ) + + return KnowledgeUploadService( + storage_dir=Path("data/files"), + milvus_client=milvus, + embedder=lambda texts: embed_texts(texts, client=llm_client), + publisher=publish_event, + document_exists=lambda doc_id: document_exists_in_milvus(milvus, doc_id), + ) diff --git a/service/customer_agent/chat.py b/service/customer_agent/chat.py new file mode 100644 index 0000000..bb5d289 --- /dev/null +++ b/service/customer_agent/chat.py @@ -0,0 +1,120 @@ +"""Anonymous customer-service orchestration without private customer access.""" +from __future__ import annotations + +import json +import re +from inspect import isawaitable + +from rag.intent import Intent + + +class QueryTooLongError(ValueError): + pass + + +async def _config(config_getter, key: str, default: str): + value = config_getter(key, default) + if isawaitable(value): + value = await value + return value or default + + +async def _maybe_await(value): + return await value if isawaitable(value) else value + + +class AnonymousCustomerAgent: + def __init__( + self, + *, + context, + rag_retrieve, + intent_recognize, + generate_answer, + audit_writer, + config_getter, + ): + self.context = context + self.rag_retrieve = rag_retrieve + self.intent_recognize = intent_recognize + self.generate_answer = generate_answer + self.audit_writer = audit_writer + self.config_getter = config_getter + + async def handle(self, session_id: str, query: str, *, trace_id: str) -> dict: + if len(query) > 2000: + raise QueryTooLongError("query长度不能超过2000字符") + await self.context.append(session_id, "user", query) + if self._contains_sensitive_input(query): + await _maybe_await(self.audit_writer( + action="anon_sensitive_input", + trace_id=trace_id, + session_id=session_id, + )) + + intent = await _maybe_await(self.intent_recognize(query)) + sources = [] + if intent is Intent.GUIDE_PURCHASE: + answer = await _config( + self.config_getter, + "agent.customer.template.guide_purchase", + "请前往开户页面办理。", + ) + elif intent is Intent.WANT_ADVISOR: + answer = await _config( + self.config_getter, + "agent.customer.template.guide_advisor", + "如需基金推荐,请联系投资顾问。", + ) + elif intent is Intent.KNOWLEDGE_QA: + try: + sources = await _maybe_await(self.rag_retrieve(query, None)) + except Exception: + sources = [] + if not sources: + answer = await _config( + self.config_getter, + "agent.customer.template.fallback_human", + "当前未找到匹配信息,请转人工客服。", + ) + else: + messages = await self.context.get(session_id) + prompt = messages + [ + { + "role": "system", + "content": "仅根据提供的知识来源回答,不得编造基金推荐。", + }, + { + "role": "system", + "content": f"知识来源:{json.dumps(sources, ensure_ascii=False)}", + }, + ] + try: + answer = await _maybe_await(self.generate_answer(prompt)) + except Exception: + answer = await _config( + self.config_getter, + "agent.customer.template.fallback_human", + "当前服务繁忙,请转人工客服。", + ) + else: + answer = await _config( + self.config_getter, + "agent.customer.template.fallback_human", + "当前未找到匹配信息,请转人工客服。", + ) + + await self.context.append(session_id, "assistant", answer) + return { + "answer": answer, + "sources": sources, + "intent": intent.value, + "trace_id": trace_id, + } + + @staticmethod + def _contains_sensitive_input(query: str) -> bool: + return bool( + re.search(r"(? str | None: + async with self.session_factory() as session: + return await self.repo_factory(session).get_value(key, default) + + async def write_audit(self, *, action: str, trace_id: str, session_id: str): + if action != "anon_sensitive_input": + raise ValueError(f"Unsupported anonymous audit action: {action}") + async with self.session_factory() as session: + await write_anonymous_sensitive_audit( + session, session_id=session_id, trace_id=trace_id + ) diff --git a/service/customer_agent/runtime.py b/service/customer_agent/runtime.py new file mode 100644 index 0000000..45b5b0a --- /dev/null +++ b/service/customer_agent/runtime.py @@ -0,0 +1,62 @@ +"""Dependency assembly for the anonymous customer-service Agent.""" +from __future__ import annotations + +from types import SimpleNamespace + +from agent.customer_agent.context import RedisConversationContext +from agent.customer_agent.session import AnonymousSessionService +from rag.embedding import embed_texts +from rag.generation import generate_answer +from rag.intent import intent_recognize +from rag.retrieve import rag_retrieve +from service.customer_agent.chat import AnonymousCustomerAgent + + +def build_anonymous_runtime( + *, + redis, + milvus_client, + llm_client, + config_getter, + audit_writer, +): + session_service = AnonymousSessionService( + redis, config_getter=config_getter + ) + context = RedisConversationContext( + redis, config_getter=config_getter + ) + + async def retrieve(query, customer_id): + return await rag_retrieve( + query, + None, + milvus_client=milvus_client, + embedder=lambda texts: embed_texts(texts, client=llm_client), + config_getter=config_getter, + ) + + async def recognize(query): + return await intent_recognize(query, llm_client=llm_client) + + async def generate(messages): + return await generate_answer( + messages, + llm_client=llm_client, + config_getter=config_getter, + ) + + agent = AnonymousCustomerAgent( + context=context, + rag_retrieve=retrieve, + intent_recognize=recognize, + generate_answer=generate, + audit_writer=audit_writer, + config_getter=config_getter, + ) + return SimpleNamespace( + redis=redis, + session_service=session_service, + context=context, + agent=agent, + ) diff --git a/service/knowledge_base/__init__.py b/service/knowledge_base/__init__.py new file mode 100644 index 0000000..a3c13be --- /dev/null +++ b/service/knowledge_base/__init__.py @@ -0,0 +1 @@ +"""Knowledge-base upload and ingestion services.""" diff --git a/service/knowledge_base/upload.py b/service/knowledge_base/upload.py new file mode 100644 index 0000000..140349a --- /dev/null +++ b/service/knowledge_base/upload.py @@ -0,0 +1,262 @@ +"""Frontend document upload, preview, and confirmed Milvus ingestion.""" +from __future__ import annotations + +import logging +import json +import uuid +from inspect import isawaitable +from pathlib import Path + +from rag.document_parser import SUPPORTED_EXTENSIONS +from rag.events import publish_knowledge_update +from rag.ingestion import ingest_document_atomic +from rag.preview import preview_document +from rag.chunk_config import resolve_chunk_config +from rag.milvus_collections import KNOWLEDGE_COLLECTIONS +from rag.milvus_delete import delete_document_vectors + + +logger = logging.getLogger("service.knowledge_base.upload") + + +class UploadValidationError(ValueError): + pass + + +class KnowledgeUploadService: + def __init__( + self, + *, + storage_dir: str | Path, + milvus_client, + embedder, + publisher, + max_upload_bytes: int = 10 * 1024 * 1024, + document_exists=None, + upload_ttl_seconds: int = 1800, + ): + self.storage_dir = Path(storage_dir) + self.storage_dir.mkdir(parents=True, exist_ok=True) + self.milvus_client = milvus_client + self.embedder = embedder + self.publisher = publisher + self.max_upload_bytes = max_upload_bytes + self.document_exists = document_exists + self.upload_ttl_seconds = upload_ttl_seconds + + async def preview( + self, + filename: str, + content: bytes, + *, + strategy: str, + chunk_size: int | None = None, + chunk_overlap: int | None = None, + ) -> dict: + self.cleanup_expired_uploads() + suffix = self._validate_upload(filename, content) + upload_id = uuid.uuid4().hex + stored_filename = f"{upload_id}{suffix}" + path = self.storage_dir / stored_filename + path.write_bytes(content) + try: + result = preview_document( + path, + strategy=strategy, + chunk_size=chunk_size, + chunk_overlap=chunk_overlap, + ) + except Exception: + path.unlink(missing_ok=True) + raise + self._manifest_path(upload_id).write_text( + json.dumps( + { + "strategy": strategy, + "chunk_size": result["chunk_size"], + "chunk_overlap": result["chunk_overlap"], + }, + ensure_ascii=False, + ), + encoding="utf-8", + ) + return { + "upload_id": upload_id, + "stored_filename": stored_filename, + "filename": Path(filename).name, + **result, + } + + def cleanup_expired_uploads(self, *, now: float | None = None) -> int: + """Remove expired upload inputs and their preview manifests.""" + if self.upload_ttl_seconds <= 0: + return 0 + if now is None: + import time + + now = time.time() + cutoff = now - self.upload_ttl_seconds + removed = 0 + for path in self.storage_dir.iterdir(): + if not path.is_file() or path.suffix.lower() not in SUPPORTED_EXTENSIONS: + continue + if path.stat().st_mtime > cutoff: + continue + upload_id = path.stem + path.unlink(missing_ok=True) + self._manifest_path(upload_id).unlink(missing_ok=True) + removed += 1 + return removed + + async def delete_document(self, doc_id: str) -> dict: + if not doc_id or not doc_id.strip(): + raise UploadValidationError("doc_id不能为空") + await delete_document_vectors(doc_id, milvus_client=self.milvus_client) + await publish_knowledge_update( + self.publisher, + doc_id=doc_id, + collection_name="", + chunk_count=0, + action="deleted", + ) + return {"doc_id": doc_id, "deleted": True} + + async def list_documents(self) -> list[dict]: + documents = {} + for collection_name in KNOWLEDGE_COLLECTIONS: + rows = await self.milvus_client.query( + collection_name=collection_name, + filter="", + output_fields=["doc_id", "title", "strategy"], + ) + for row in rows or []: + doc_id = row.get("doc_id", "") + if not doc_id: + continue + document = documents.setdefault( + doc_id, + { + "doc_id": doc_id, + "title": row.get("title", ""), + "collection_name": collection_name, + "strategy": row.get("strategy", ""), + "chunk_count": 0, + }, + ) + document["chunk_count"] += 1 + return list(documents.values()) + + async def get_document(self, doc_id: str) -> dict | None: + if not doc_id or not doc_id.strip(): + raise UploadValidationError("doc_id不能为空") + escaped = doc_id.replace("\\", "\\\\").replace('"', '\\"') + for collection_name in KNOWLEDGE_COLLECTIONS: + rows = await self.milvus_client.query( + collection_name=collection_name, + filter=f'doc_id == "{escaped}"', + output_fields=["doc_id", "title", "section_title", "strategy", "text"], + ) + if rows: + first = rows[0] + return { + "doc_id": first.get("doc_id", doc_id), + "title": first.get("title", ""), + "collection_name": collection_name, + "strategy": first.get("strategy", ""), + "chunk_count": len(rows), + "chunks": [ + { + "section_title": row.get("section_title") or None, + "text": row.get("text", ""), + } + for row in rows + ], + } + return None + + async def confirm( + self, + *, + upload_id: str, + title: str, + doc_id: str, + collection_name: str, + strategy: str, + chunk_size: int | None = None, + chunk_overlap: int | None = None, + ) -> dict: + self.cleanup_expired_uploads() + if collection_name not in KNOWLEDGE_COLLECTIONS: + raise UploadValidationError("知识库集合不在允许范围内") + path = self._find_upload(upload_id) + try: + manifest = json.loads(self._manifest_path(upload_id).read_text(encoding="utf-8")) + if strategy != manifest["strategy"]: + raise UploadValidationError("确认策略必须与预览策略一致") + if chunk_size is None: + chunk_size = manifest["chunk_size"] + if chunk_overlap is None: + chunk_overlap = manifest["chunk_overlap"] + if self.document_exists is not None: + exists = self.document_exists(doc_id) + if isawaitable(exists): + exists = await exists + if exists: + raise UploadValidationError("doc_id已存在,禁止重复入库") + from rag.document_parser import parse_document + + result = await ingest_document_atomic( + parse_document(path), + doc_id, + title, + collection_name, + strategy, + milvus_client=self.milvus_client, + embedder=self.embedder, + config=resolve_chunk_config( + chunk_size=chunk_size, + chunk_overlap=chunk_overlap, + ), + ) + try: + await publish_knowledge_update( + self.publisher, + doc_id=doc_id, + collection_name=collection_name, + chunk_count=result["chunk_count"], + ) + except Exception: + logger.exception("knowledge update event publish failed for doc_id=%s", doc_id) + try: + await delete_document_vectors(doc_id, milvus_client=self.milvus_client) + except Exception: + logger.exception("failed to compensate vectors for doc_id=%s", doc_id) + raise + return {"event_published": True, **result} + finally: + path.unlink(missing_ok=True) + self._manifest_path(upload_id).unlink(missing_ok=True) + + def _validate_upload(self, filename: str, content: bytes) -> str: + suffix = Path(filename).suffix.lower() + if suffix not in SUPPORTED_EXTENSIONS: + raise UploadValidationError(f"Unsupported document type: {suffix or ''}") + if not content: + raise UploadValidationError("上传文件不能为空") + if len(content) > self.max_upload_bytes: + raise UploadValidationError("上传文件超过大小限制") + return suffix + + def _find_upload(self, upload_id: str) -> Path: + if not upload_id or Path(upload_id).name != upload_id: + raise UploadValidationError("upload_id无效") + paths = [ + path for path in self.storage_dir.glob(f"{upload_id}.*") + if path.suffix.lower() in SUPPORTED_EXTENSIONS + ] + if len(paths) != 1: + raise UploadValidationError("上传文件不存在或已过期") + return paths[0] + + def _manifest_path(self, upload_id: str) -> Path: + return self.storage_dir / f"{upload_id}.json" diff --git a/tests/test_agent_context.py b/tests/test_agent_context.py new file mode 100644 index 0000000..2bf997c --- /dev/null +++ b/tests/test_agent_context.py @@ -0,0 +1,56 @@ +import json +import unittest + +from agent.customer_agent.context import RedisConversationContext + + +class ContextRedis: + def __init__(self): + self.values = {} + + async def rpush(self, key, value): + self.values.setdefault(key, []).append(value) + + async def lrange(self, key, start, end): + values = self.values.get(key, [])[start:] + return values if end == -1 else values[: end - start + 1] + + async def ltrim(self, key, start, end): + values = self.values.get(key, [])[start:] + self.values[key] = values if end == -1 else values[: end - start + 1] + + async def expire(self, key, seconds): + pass + + +class ContextTests(unittest.IsolatedAsyncioTestCase): + async def test_appends_messages_and_trims_oldest_when_token_limit_is_exceeded(self): + redis = ContextRedis() + context = RedisConversationContext( + redis, + config_getter={ + "agent.customer.session_max_token": "5", + "agent.customer.session.ttl": "120", + }.get, + token_counter=lambda text: len(text), + ) + + await context.append("s1", "user", "123") + await context.append("s1", "assistant", "456") + await context.append("s1", "user", "789") + + messages = await context.get("s1") + self.assertEqual([item["content"] for item in messages], ["789"]) + + async def test_message_payload_is_structured_json(self): + redis = ContextRedis() + context = RedisConversationContext(redis, config_getter={}.get) + + await context.append("s1", "user", "基金是什么") + + raw = redis.values["session:s1:messages"][0] + self.assertEqual(json.loads(raw)["role"], "user") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_agent_session.py b/tests/test_agent_session.py new file mode 100644 index 0000000..f8f9d8e --- /dev/null +++ b/tests/test_agent_session.py @@ -0,0 +1,88 @@ +import unittest + +from agent.customer_agent.session import AnonymousSessionService, SessionOwnershipError + + +class FakeRedis: + def __init__(self): + self.values = {} + self.lists = {} + self.sorted_sets = {} + self.expirations = {} + + async def set(self, key, value, ex=None): + self.values[key] = value + if ex is not None: + self.expirations[key] = ex + + async def exists(self, key): + return int(key in self.values or key in self.lists) + + async def rpush(self, key, value): + self.lists.setdefault(key, []).append(value) + return len(self.lists[key]) + + async def expire(self, key, seconds): + self.expirations[key] = seconds + + async def zremrangebyscore(self, key, minimum, maximum): + self.sorted_sets[key] = [item for item in self.sorted_sets.get(key, []) if item[0] > maximum] + + async def zadd(self, key, mapping): + self.sorted_sets.setdefault(key, []).extend((score, member) for member, score in mapping.items()) + + async def zcard(self, key): + return len(self.sorted_sets.get(key, [])) + + +class SessionTests(unittest.IsolatedAsyncioTestCase): + async def test_create_initializes_anonymous_session_and_message_list(self): + redis = FakeRedis() + service = AnonymousSessionService(redis, config_getter={"agent.customer.session.ttl": "120"}.get) + + session_id = await service.create_session() + + self.assertEqual(len(session_id), 32) + self.assertIn(f"session:{session_id}", redis.values) + self.assertIn(f"session:{session_id}:messages", redis.lists) + self.assertEqual(redis.expirations[f"session:{session_id}"], 120) + + async def test_missing_session_raises_not_found(self): + service = AnonymousSessionService(FakeRedis(), config_getter={}.get) + + with self.assertRaises(SessionOwnershipError) as ctx: + await service.verify_session_ownership("missing") + + self.assertEqual(ctx.exception.code, 404) + + async def test_authenticated_owner_mismatch_raises_forbidden(self): + redis = FakeRedis() + service = AnonymousSessionService(redis, config_getter={}.get) + session_id = await service.create_session() + + with self.assertRaises(SessionOwnershipError) as ctx: + await service.verify_session_ownership(session_id, customer_id="customer-1", session_customer_id="customer-2") + + self.assertEqual(ctx.exception.code, 403) + + async def test_rate_limit_uses_sliding_window_and_returns_retry_after(self): + redis = FakeRedis() + service = AnonymousSessionService( + redis, + config_getter={ + "agent.customer.rate_limit.window_sec": "60", + "agent.customer.rate_limit.max_requests": "1", + "agent.customer.session.ttl": "120", + }.get, + clock=lambda: 100.0, + ) + await service.create_session() + + self.assertEqual(await service.consume_chat_quota("s"), None) + retry_after = await service.consume_chat_quota("s") + + self.assertEqual(retry_after, 60) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_anonymous_agent_service.py b/tests/test_anonymous_agent_service.py new file mode 100644 index 0000000..e6c727e --- /dev/null +++ b/tests/test_anonymous_agent_service.py @@ -0,0 +1,94 @@ +import unittest +from unittest.mock import AsyncMock + +from rag.intent import Intent +from service.customer_agent.chat import AnonymousCustomerAgent, QueryTooLongError + + +class AnonymousAgentTests(unittest.IsolatedAsyncioTestCase): + def _build(self, intent, sources=None, answer="LLM answer"): + context = AsyncMock() + rag = AsyncMock(return_value=sources or []) + recognize = AsyncMock(return_value=intent) + generate = AsyncMock(return_value=answer) + audit = AsyncMock() + config = { + "agent.customer.template.guide_purchase": "请前往开户页面办理", + "agent.customer.template.guide_advisor": "如需推荐,请联系投顾", + "agent.customer.template.fallback_human": "请转人工客服", + } + service = AnonymousCustomerAgent( + context=context, + rag_retrieve=rag, + intent_recognize=recognize, + generate_answer=generate, + audit_writer=audit, + config_getter=config.get, + ) + return service, context, rag, recognize, generate, audit + + async def test_knowledge_question_uses_rag_without_customer_memory(self): + service, context, rag, _, generate, _ = self._build( + Intent.KNOWLEDGE_QA, + sources=[{"doc_id": "d1", "chunk_text": "基金正文", "score": 0.9}], + ) + + result = await service.handle("s1", "基金是什么", trace_id="trace-1") + + rag.assert_awaited_once_with("基金是什么", None) + generate.assert_awaited_once() + self.assertEqual(result["sources"][0]["doc_id"], "d1") + self.assertEqual(result["trace_id"], "trace-1") + self.assertEqual(context.append.await_count, 2) + + async def test_recommendation_returns_fixed_advisor_guidance_without_llm(self): + service, _, rag, _, generate, _ = self._build(Intent.WANT_ADVISOR) + + result = await service.handle("s1", "给我推荐一只基金", trace_id="trace-2") + + self.assertEqual(result["answer"], "如需推荐,请联系投顾") + rag.assert_not_awaited() + generate.assert_not_awaited() + self.assertEqual(result["sources"], []) + + async def test_sensitive_input_is_audited_and_no_private_data_is_loaded(self): + service, _, rag, _, generate, audit = self._build(Intent.KNOWLEDGE_QA, sources=[]) + + await service.handle("s1", "我的手机号是13812345678", trace_id="trace-3") + + audit.assert_awaited_once() + self.assertEqual(audit.await_args.kwargs["action"], "anon_sensitive_input") + self.assertEqual(audit.await_args.kwargs["trace_id"], "trace-3") + rag.assert_awaited_once_with("我的手机号是13812345678", None) + generate.assert_not_awaited() + + async def test_query_over_2000_characters_is_rejected(self): + service, *_ = self._build(Intent.KNOWLEDGE_QA) + + with self.assertRaises(QueryTooLongError): + await service.handle("s1", "x" * 2001, trace_id="trace-4") + + async def test_milvus_failure_returns_human_fallback(self): + service, _, rag, _, generate, _ = self._build(Intent.KNOWLEDGE_QA) + rag.side_effect = TimeoutError("Milvus down") + + result = await service.handle("s1", "基金是什么", trace_id="trace-5") + + self.assertEqual(result["answer"], "请转人工客服") + generate.assert_not_awaited() + + async def test_llm_failure_returns_human_fallback_without_fake_sources(self): + service, _, _, _, generate, _ = self._build( + Intent.KNOWLEDGE_QA, + sources=[{"doc_id": "d1", "chunk_text": "正文", "score": 0.8}], + ) + generate.side_effect = RuntimeError("LLM down") + + result = await service.handle("s1", "基金是什么", trace_id="trace-6") + + self.assertEqual(result["answer"], "请转人工客服") + self.assertEqual(result["sources"][0]["doc_id"], "d1") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_chunk_config.py b/tests/test_chunk_config.py new file mode 100644 index 0000000..a7841b8 --- /dev/null +++ b/tests/test_chunk_config.py @@ -0,0 +1,29 @@ +import unittest + +from rag.chunk_config import DEFAULT_CHUNK_OVERLAP, DEFAULT_CHUNK_SIZE, resolve_chunk_config + + +class ChunkConfigTests(unittest.TestCase): + def test_uses_project_defaults_without_sys_config(self): + config = resolve_chunk_config() + + self.assertEqual(config.size, DEFAULT_CHUNK_SIZE) + self.assertEqual(config.overlap, DEFAULT_CHUNK_OVERLAP) + self.assertEqual((config.size, config.overlap), (512, 64)) + + def test_custom_values_override_defaults_independently(self): + config = resolve_chunk_config(chunk_size=256) + + self.assertEqual((config.size, config.overlap), (256, 64)) + + def test_rejects_invalid_chunk_values(self): + with self.assertRaises(ValueError): + resolve_chunk_config(chunk_size=0) + with self.assertRaises(ValueError): + resolve_chunk_config(chunk_size=64, chunk_overlap=64) + with self.assertRaises(ValueError): + resolve_chunk_config(chunk_overlap=-1) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_chunking.py b/tests/test_chunking.py new file mode 100644 index 0000000..017861d --- /dev/null +++ b/tests/test_chunking.py @@ -0,0 +1,88 @@ +import unittest + +from rag.chunk_config import ChunkConfig +from rag.chunking import chunk_document + + +class DefaultChunkingTests(unittest.TestCase): + def test_keeps_short_paragraphs_as_separate_chunks(self): + result = chunk_document( + "第一段内容。\n\n第二段内容。", + "default", + config=ChunkConfig(size=64, overlap=4), + ) + + self.assertEqual([chunk.text for chunk in result.chunks], ["第一段内容。", "第二段内容。"]) + self.assertEqual(result.actual_strategy, "default") + self.assertFalse(result.degraded) + + def test_splits_long_paragraphs_with_overlap(self): + result = chunk_document( + "abcdefghij" * 3, + "default", + config=ChunkConfig(size=10, overlap=2), + ) + + self.assertGreater(len(result.chunks), 1) + self.assertTrue(all(len(chunk.text) <= 10 for chunk in result.chunks)) + self.assertEqual(result.chunks[0].text[-2:], result.chunks[1].text[:2]) + + +class QaPairChunkingTests(unittest.TestCase): + def test_keeps_each_question_and_answer_together(self): + result = chunk_document( + "Q: 什么是净值?\nA: 净值是基金单位价值。\nQ: 如何申购?\nA: 通过交易页面申购。", + "qa_pair", + config=ChunkConfig(size=64, overlap=4), + ) + + self.assertEqual(len(result.chunks), 2) + self.assertIn("question: 什么是净值?", result.chunks[0].text) + self.assertIn("answer: 净值是基金单位价值。", result.chunks[0].text) + + def test_rejects_missing_qa_markers(self): + with self.assertRaises(ValueError): + chunk_document("普通文本,没有问答标记", "qa_pair") + + def test_rejects_the_whole_document_when_one_pair_is_too_long(self): + text = "Q: 短问题\nA: 短答案\nQ: 长问题\nA: " + ("很长" * 20) + + with self.assertRaises(ValueError): + chunk_document(text, "qa_pair", config=ChunkConfig(size=20, overlap=2)) + + +class ChapterChunkingTests(unittest.TestCase): + def test_preserves_nested_heading_path(self): + result = chunk_document( + "# 产品说明\n## 风险揭示\n风险内容。", + "chapter_semantic", + config=ChunkConfig(size=64, overlap=4), + ) + + self.assertEqual(result.actual_strategy, "chapter_semantic") + self.assertEqual(result.chunks[0].section_title, "产品说明 > 风险揭示") + self.assertTrue(result.chunks[0].text.startswith("【章节:产品说明 > 风险揭示】")) + + def test_splits_only_inside_a_chapter(self): + result = chunk_document( + "# 第一章\n" + ("甲" * 20) + "\n# 第二章\n" + ("乙" * 20), + "chapter_semantic", + config=ChunkConfig(size=16, overlap=2), + ) + + self.assertTrue(all("甲" not in chunk.text or "乙" not in chunk.text for chunk in result.chunks)) + + def test_degrades_to_default_without_headings(self): + result = chunk_document( + "没有Markdown标题的普通内容。", + "chapter_semantic", + config=ChunkConfig(size=64, overlap=4), + ) + + self.assertEqual(result.actual_strategy, "default") + self.assertTrue(result.degraded) + self.assertTrue(result.warning) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_cleaning.py b/tests/test_cleaning.py new file mode 100644 index 0000000..610c703 --- /dev/null +++ b/tests/test_cleaning.py @@ -0,0 +1,30 @@ +import unittest + +from rag.cleaning import clean_document_text + + +class CleaningTests(unittest.TestCase): + def test_normalizes_bom_line_endings_whitespace_and_blank_lines(self): + result = clean_document_text("\ufeff标题\r\n\r\n\r\n段落\t 内容\x00") + + self.assertEqual(result.text, "标题\n\n段落 内容") + self.assertTrue(result.changed) + self.assertTrue(result.warnings) + + def test_removes_standalone_page_markers_but_keeps_business_numbers(self): + result = clean_document_text("产品代码 110011\n第 1 页\n年化收益率 3.5%\nPage 2") + + self.assertEqual(result.text, "产品代码 110011\n年化收益率 3.5%") + self.assertIn("110011", result.text) + + def test_does_not_remove_meaningful_content(self): + text = "基金名称:示例基金\n风险等级:R3\n申购费率:1.2%" + + result = clean_document_text(text) + + self.assertEqual(result.text, text) + self.assertFalse(result.changed) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_customer_agent_audit.py b/tests/test_customer_agent_audit.py new file mode 100644 index 0000000..2924533 --- /dev/null +++ b/tests/test_customer_agent_audit.py @@ -0,0 +1,24 @@ +import unittest +from unittest.mock import AsyncMock + +from service.customer_agent.audit import write_anonymous_sensitive_audit + + +class AuditTests(unittest.IsolatedAsyncioTestCase): + async def test_writes_sensitive_anonymous_action_with_trace_id(self): + db = AsyncMock() + + await write_anonymous_sensitive_audit( + db, session_id="s1", trace_id="trace-1" + ) + + db.execute.assert_awaited_once() + db.commit.assert_awaited_once() + params = db.execute.await_args.args[1] + self.assertEqual(params["action"], "anon_sensitive_input") + self.assertEqual(params["trace_id"], "trace-1") + self.assertEqual(params["target"], "s1") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_customer_agent_bootstrap.py b/tests/test_customer_agent_bootstrap.py new file mode 100644 index 0000000..d6b41cd --- /dev/null +++ b/tests/test_customer_agent_bootstrap.py @@ -0,0 +1,42 @@ +import asyncio +import unittest +from unittest.mock import patch + +from service.customer_agent.bootstrap import build_default_runtime, document_exists_in_milvus + + +class BootstrapTests(unittest.TestCase): + def test_document_exists_checks_all_knowledge_collections(self): + class FakeMilvus: + def __init__(self): + self.calls = [] + + async def query(self, **kwargs): + self.calls.append(kwargs) + return [{"doc_id": "doc-1"}] if kwargs["collection_name"] == "fin_policy" else [] + + async def run(): + client = FakeMilvus() + self.assertTrue(await document_exists_in_milvus(client, "doc-1")) + self.assertEqual( + [call["collection_name"] for call in client.calls], + ["fin_faq", "fin_fund_doc", "fin_policy"], + ) + + asyncio.run(run()) + + def test_builds_runtime_from_project_clients_and_mysql_config(self): + with patch("service.customer_agent.bootstrap.redis_client", return_value="redis"), \ + patch("service.customer_agent.bootstrap.milvus_client", return_value="milvus"), \ + patch("service.customer_agent.bootstrap.llm_client", new="llm"), \ + patch("service.customer_agent.bootstrap.build_anonymous_runtime", return_value="runtime") as builder: + result = build_default_runtime() + + self.assertEqual(result, "runtime") + self.assertEqual(builder.call_args.kwargs["redis"], "redis") + self.assertEqual(builder.call_args.kwargs["milvus_client"], "milvus") + self.assertEqual(builder.call_args.kwargs["llm_client"], "llm") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_customer_agent_config.py b/tests/test_customer_agent_config.py new file mode 100644 index 0000000..4c69d62 --- /dev/null +++ b/tests/test_customer_agent_config.py @@ -0,0 +1,30 @@ +import unittest +from unittest.mock import AsyncMock + +from service.customer_agent.config import DatabaseConfigProvider + + +class ConfigProviderTests(unittest.IsolatedAsyncioTestCase): + async def test_reads_sys_config_value_and_uses_default_when_missing(self): + repo = AsyncMock() + repo.get_value.side_effect = ["120", "fallback"] + + class SessionContext: + async def __aenter__(self): + return object() + + async def __aexit__(self, *args): + pass + + provider = DatabaseConfigProvider( + repo_factory=lambda session: repo, + session_factory=lambda: SessionContext(), + ) + + self.assertEqual(await provider.get("agent.customer.session.ttl", "60"), "120") + self.assertEqual(await provider.get("missing", "fallback"), "fallback") + self.assertEqual(repo.get_value.await_count, 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_customer_agent_packages.py b/tests/test_customer_agent_packages.py new file mode 100644 index 0000000..c09d2e9 --- /dev/null +++ b/tests/test_customer_agent_packages.py @@ -0,0 +1,14 @@ +import unittest + + +class CustomerAgentPackageTests(unittest.TestCase): + def test_agent_and_service_customer_agent_packages_are_importable(self): + from agent.customer_agent import package_name as agent_package_name + from service.customer_agent import package_name as service_package_name + + self.assertEqual(agent_package_name, "customer_agent") + self.assertEqual(service_package_name, "customer_agent") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_customer_agent_router.py b/tests/test_customer_agent_router.py new file mode 100644 index 0000000..2604a94 --- /dev/null +++ b/tests/test_customer_agent_router.py @@ -0,0 +1,99 @@ +import json +import unittest +from types import SimpleNamespace + +import httpx +from fastapi import FastAPI + +from api.routers.customer_agent import router +from rag.intent import Intent +from service.customer_agent.chat import AnonymousCustomerAgent +from agent.customer_agent.session import AnonymousSessionService +from tests.test_agent_session import FakeRedis + + +class RouterTests(unittest.IsolatedAsyncioTestCase): + async def asyncSetUp(self): + redis = FakeRedis() + config = { + "agent.customer.session.ttl": "120", + "agent.customer.rate_limit.window_sec": "60", + "agent.customer.rate_limit.max_requests": "1", + "agent.customer.template.fallback_human": "请转人工客服", + } + session = AnonymousSessionService(redis, config_getter=config.get) + + class Context: + async def append(self, *args): + pass + + async def get(self, *args): + return [] + + agent = AnonymousCustomerAgent( + context=Context(), + rag_retrieve=lambda query, customer_id: [], + intent_recognize=lambda query: Intent.NO_MATCH, + generate_answer=lambda messages: "answer", + audit_writer=lambda **kwargs: None, + config_getter=config.get, + ) + self.app = FastAPI() + self.app.state.customer_agent_runtime = SimpleNamespace( + redis=redis, session_service=session, agent=agent + ) + self.app.include_router(router, prefix="/api/agent/customer") + + async def test_create_chat_and_end_anonymous_session(self): + transport = httpx.ASGITransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + created = await client.post("/api/agent/customer/session/create") + self.assertEqual(created.status_code, 200) + session_id = created.json()["data"]["session_id"] + + chat = await client.get( + "/api/agent/customer/chat", + params={"session_id": session_id, "query": "你好"}, + headers={"X-Trace-Id": "trace-router"}, + ) + self.assertEqual(chat.status_code, 200) + event = json.loads(chat.text.removeprefix("data: ").strip()) + self.assertEqual(event["trace_id"], "trace-router") + + limited = await client.get( + "/api/agent/customer/chat", + params={"session_id": session_id, "query": "再次提问"}, + ) + self.assertEqual(limited.status_code, 429) + self.assertIn(int(limited.headers["Retry-After"]), range(1, 61)) + + ended = await client.post( + "/api/agent/customer/session/end", json={"session_id": session_id} + ) + self.assertEqual(ended.status_code, 200) + + async def test_query_over_2000_returns_400(self): + transport = httpx.ASGITransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + created = await client.post("/api/agent/customer/session/create") + session_id = created.json()["data"]["session_id"] + response = await client.get( + "/api/agent/customer/chat", + params={"session_id": session_id, "query": "x" * 2001}, + ) + self.assertEqual(response.status_code, 400) + self.assertEqual(response.json()["code"], 400) + + async def test_missing_session_returns_404(self): + transport = httpx.ASGITransport(app=self.app) + async with httpx.AsyncClient(transport=transport, base_url="http://test") as client: + response = await client.get( + "/api/agent/customer/chat", + params={"session_id": "missing", "query": "你好"}, + ) + self.assertEqual(response.status_code, 404) + self.assertEqual(response.json()["code"], 404) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_customer_agent_runtime.py b/tests/test_customer_agent_runtime.py new file mode 100644 index 0000000..2d1934b --- /dev/null +++ b/tests/test_customer_agent_runtime.py @@ -0,0 +1,25 @@ +import unittest + +from service.customer_agent.runtime import build_anonymous_runtime +from agent.customer_agent.session import AnonymousSessionService +from agent.customer_agent.context import RedisConversationContext + + +class RuntimeTests(unittest.TestCase): + def test_runtime_wires_session_context_and_agent_components(self): + runtime = build_anonymous_runtime( + redis=object(), + milvus_client=object(), + llm_client=object(), + config_getter={}.get, + audit_writer=lambda **kwargs: None, + ) + + self.assertIsInstance(runtime.session_service, AnonymousSessionService) + self.assertIsInstance(runtime.context, RedisConversationContext) + self.assertIs(runtime.redis, runtime.session_service.redis) + self.assertIsNotNone(runtime.agent) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_document_parser.py b/tests/test_document_parser.py new file mode 100644 index 0000000..352eb98 --- /dev/null +++ b/tests/test_document_parser.py @@ -0,0 +1,72 @@ +import tempfile +import unittest +from pathlib import Path +from zipfile import ZIP_DEFLATED, ZipFile + +from pypdf import PdfWriter +from pypdf.generic import DecodedStreamObject, DictionaryObject, NameObject + +from rag.document_parser import parse_document + + +class DocumentParserTests(unittest.TestCase): + def test_parses_txt_and_md_as_utf8_text(self): + with tempfile.TemporaryDirectory() as tmp: + root = Path(tmp) + txt = root / "notice.txt" + md = root / "notice.md" + txt.write_text("基金风险提示", encoding="utf-8") + md.write_text("# 产品说明\n\n开放式基金", encoding="utf-8") + + self.assertEqual(parse_document(txt), "基金风险提示") + self.assertEqual(parse_document(md), "# 产品说明\n\n开放式基金") + + def test_extracts_text_from_docx(self): + with tempfile.TemporaryDirectory() as tmp: + docx = Path(tmp) / "notice.docx" + document_xml = ( + '' + '' + '基金产品说明' + '风险揭示' + ) + with ZipFile(docx, "w", ZIP_DEFLATED) as archive: + archive.writestr("word/document.xml", document_xml) + + self.assertEqual(parse_document(docx), "基金产品说明\n风险揭示") + + def test_extracts_text_from_pdf(self): + with tempfile.TemporaryDirectory() as tmp: + pdf = Path(tmp) / "notice.pdf" + writer = PdfWriter() + page = writer.add_blank_page(width=612, height=792) + font = writer._add_object( + DictionaryObject( + { + NameObject("/Type"): NameObject("/Font"), + NameObject("/Subtype"): NameObject("/Type1"), + NameObject("/BaseFont"): NameObject("/Helvetica"), + } + ) + ) + page[NameObject("/Resources")] = DictionaryObject( + {NameObject("/Font"): DictionaryObject({NameObject("/F1"): font})} + ) + page[NameObject("/Contents")] = DecodedStreamObject() + page[NameObject("/Contents")].set_data(b"BT /F1 12 Tf 72 720 Td (Fund FAQ) Tj ET") + with pdf.open("wb") as output: + writer.write(output) + + self.assertEqual(parse_document(pdf), "Fund FAQ") + + def test_rejects_unsupported_extension(self): + with tempfile.TemporaryDirectory() as tmp: + path = Path(tmp) / "notice.xlsx" + path.write_bytes(b"not supported") + + with self.assertRaises(ValueError): + parse_document(path) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_embedding.py b/tests/test_embedding.py new file mode 100644 index 0000000..17c9a0e --- /dev/null +++ b/tests/test_embedding.py @@ -0,0 +1,35 @@ +import unittest + +from rag.embedding import EmbeddingError, embed_texts + + +class EmbeddingTests(unittest.IsolatedAsyncioTestCase): + async def test_returns_768_dimension_vectors(self): + class Client: + async def embed(self, texts): + return [[0.1] * 768 for _ in texts] + + vectors = await embed_texts(["基金知识"], client=Client()) + + self.assertEqual(len(vectors), 1) + self.assertEqual(len(vectors[0]), 768) + + async def test_rejects_wrong_embedding_dimension(self): + class Client: + async def embed(self, texts): + return [[0.1] * 3 for _ in texts] + + with self.assertRaises(EmbeddingError): + await embed_texts(["基金知识"], client=Client()) + + async def test_wraps_provider_failure(self): + class Client: + async def embed(self, texts): + raise TimeoutError("embedding timeout") + + with self.assertRaises(EmbeddingError): + await embed_texts(["基金知识"], client=Client()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_events.py b/tests/test_events.py new file mode 100644 index 0000000..eef92e3 --- /dev/null +++ b/tests/test_events.py @@ -0,0 +1,22 @@ +import unittest +from unittest.mock import AsyncMock + +from rag.events import publish_knowledge_update + + +class KnowledgeEventTests(unittest.IsolatedAsyncioTestCase): + async def test_publishes_update_event_after_successful_ingestion(self): + publisher = AsyncMock() + + await publish_knowledge_update( + publisher, doc_id="doc-1", collection_name="fin_faq", chunk_count=4 + ) + + publisher.assert_awaited_once_with( + "event:knowledge_update", + {"doc_id": "doc-1", "collection_name": "fin_faq", "chunk_count": 4}, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_generation.py b/tests/test_generation.py new file mode 100644 index 0000000..8217c2e --- /dev/null +++ b/tests/test_generation.py @@ -0,0 +1,41 @@ +import unittest +from unittest.mock import AsyncMock + +from rag.generation import generate_answer + + +class GenerationTests(unittest.IsolatedAsyncioTestCase): + async def test_switches_to_backup_model_after_primary_failure(self): + llm = AsyncMock() + llm.chat.side_effect = [RuntimeError("primary down"), "backup answer"] + config = { + "agent.customer.llm.fallback_model": "backup-model", + "agent.customer.template.system_busy": "系统繁忙,请稍后再试", + } + + result = await generate_answer( + [{"role": "user", "content": "基金是什么"}], + llm_client=llm, config_getter=config.get, primary_model="primary-model", + ) + + self.assertEqual(result, "backup answer") + self.assertEqual(llm.chat.await_args_list[1].kwargs["model"], "backup-model") + + async def test_returns_configured_system_busy_template_when_models_fail(self): + llm = AsyncMock() + llm.chat.side_effect = RuntimeError("down") + config = { + "agent.customer.llm.fallback_model": "backup-model", + "agent.customer.template.system_busy": "系统繁忙,请稍后再试", + } + + result = await generate_answer( + [{"role": "user", "content": "基金是什么"}], + llm_client=llm, config_getter=config.get, primary_model="primary-model", + ) + + self.assertEqual(result, "系统繁忙,请稍后再试") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_ingestion.py b/tests/test_ingestion.py new file mode 100644 index 0000000..8e1a51d --- /dev/null +++ b/tests/test_ingestion.py @@ -0,0 +1,84 @@ +import unittest +from unittest.mock import AsyncMock + +from rag.ingestion import ingest_document_atomic + + +class AtomicIngestionTests(unittest.IsolatedAsyncioTestCase): + async def test_validation_completes_before_embedding(self): + milvus = AsyncMock() + embedder = AsyncMock() + + with self.assertRaises(ValueError): + await ingest_document_atomic( + "Q: only a question", "doc-1", "FAQ", "fin_faq", "qa_pair", + milvus_client=milvus, embedder=embedder, + ) + + embedder.assert_not_awaited() + milvus.insert.assert_not_awaited() + + async def test_embedding_failure_cleans_up_document_rows(self): + milvus = AsyncMock() + embedder = AsyncMock(side_effect=RuntimeError("embedding down")) + + with self.assertRaises(RuntimeError): + await ingest_document_atomic( + "plain text", "doc-2", "Policy", "fin_policy", "default", + milvus_client=milvus, embedder=embedder, + ) + + milvus.delete.assert_awaited_once_with( + collection_name="fin_policy", filter='doc_id == "doc-2"' + ) + + async def test_milvus_failure_cleans_up_partial_document_rows(self): + milvus = AsyncMock() + milvus.insert.side_effect = RuntimeError("insert down") + embedder = AsyncMock(return_value=[[0.0] * 768]) + + with self.assertRaises(RuntimeError): + await ingest_document_atomic( + "plain text", "doc-3", "Policy", "fin_policy", "default", + milvus_client=milvus, embedder=embedder, + ) + + milvus.delete.assert_awaited_once_with( + collection_name="fin_policy", filter='doc_id == "doc-3"' + ) + + async def test_success_inserts_complete_document_rows(self): + milvus = AsyncMock() + embedder = AsyncMock(return_value=[[0.0] * 768]) + + result = await ingest_document_atomic( + "plain text", "doc-4", "Policy", "fin_policy", "default", + milvus_client=milvus, embedder=embedder, + ) + + self.assertEqual(result["doc_id"], "doc-4") + milvus.delete.assert_not_awaited() + milvus.insert.assert_awaited_once() + + async def test_cleans_text_before_embedding_and_milvus_insert(self): + milvus = AsyncMock() + embedder = AsyncMock(return_value=[[0.0] * 768, [0.0] * 768]) + + await ingest_document_atomic( + "\ufeff基金名称:示例基金\r\n\r\n\r\n风险等级:R3\t ", + "doc-clean", + "示例基金", + "fin_fund_doc", + "default", + milvus_client=milvus, + embedder=embedder, + ) + + self.assertEqual( + embedder.await_args.args[0], + ["基金名称:示例基金", "风险等级:R3"], + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_intent.py b/tests/test_intent.py new file mode 100644 index 0000000..1123767 --- /dev/null +++ b/tests/test_intent.py @@ -0,0 +1,34 @@ +import unittest +from unittest.mock import AsyncMock + +from rag.intent import Intent, intent_recognize + + +class IntentTests(unittest.IsolatedAsyncioTestCase): + async def test_returns_a_declared_intent_enum(self): + llm = AsyncMock() + llm.chat.return_value = "knowledge_qa" + + result = await intent_recognize("基金是什么", llm_client=llm) + + self.assertIs(result, Intent.KNOWLEDGE_QA) + + async def test_invalid_llm_output_returns_no_match(self): + llm = AsyncMock() + llm.chat.return_value = "made_up_intent" + + result = await intent_recognize("随便聊聊", llm_client=llm) + + self.assertIs(result, Intent.NO_MATCH) + + async def test_llm_failure_returns_no_match_for_handoff(self): + llm = AsyncMock() + llm.chat.side_effect = RuntimeError("LLM down") + + result = await intent_recognize("我要投诉", llm_client=llm) + + self.assertIs(result, Intent.NO_MATCH) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_knowledge_auth.py b/tests/test_knowledge_auth.py new file mode 100644 index 0000000..c50fbc2 --- /dev/null +++ b/tests/test_knowledge_auth.py @@ -0,0 +1,28 @@ +import unittest +from types import SimpleNamespace + +from api.deps import require_knowledge_operator +from utils.exceptions import ForbiddenError + + +class KnowledgeAuthTests(unittest.IsolatedAsyncioTestCase): + async def test_allows_admin_or_knowledge_operator(self): + admin = SimpleNamespace(user_type="ADMIN", employee_role=None) + operator = SimpleNamespace(user_type="EMPLOYEE", employee_role="KNOWLEDGE_OPERATOR") + + self.assertIs(await require_knowledge_operator(admin), admin) + self.assertIs(await require_knowledge_operator(operator), operator) + + async def test_rejects_customer_and_unrelated_employee(self): + with self.assertRaises(ForbiddenError): + await require_knowledge_operator( + SimpleNamespace(user_type="CUSTOMER", employee_role=None) + ) + with self.assertRaises(ForbiddenError): + await require_knowledge_operator( + SimpleNamespace(user_type="EMPLOYEE", employee_role="ADVISOR") + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_knowledge_router.py b/tests/test_knowledge_router.py new file mode 100644 index 0000000..bc7e8d5 --- /dev/null +++ b/tests/test_knowledge_router.py @@ -0,0 +1,148 @@ +import tempfile +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock + +from api.routers.knowledge import ( + confirm_document_upload, + delete_document, + get_document, + list_documents, + preview_document_upload, +) +from service.knowledge_base.upload import KnowledgeUploadService +from types import SimpleNamespace + + +class FakeUpload: + def __init__(self, filename, content): + self.filename = filename + self._content = content + + async def read(self): + return self._content + + +class FakeRequest: + def __init__(self, service, form_data=None, json_data=None): + self.app = SimpleNamespace(state=SimpleNamespace(knowledge_upload_service=service)) + self.form_data = form_data + self.json_data = json_data + + async def form(self): + return self.form_data + + async def json(self): + return self.json_data + + +class KnowledgeRouterTests(unittest.IsolatedAsyncioTestCase): + async def test_preview_endpoint_accepts_file_and_manual_strategy(self): + with tempfile.TemporaryDirectory() as tmp: + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=AsyncMock(), + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + request = FakeRequest( + service, + { + "file": FakeUpload("faq.md", b"Q: Q\nA: A"), + "strategy": "qa_pair", + }, + ) + + response = await preview_document_upload( + request, + SimpleNamespace(user_type="ADMIN", employee_role=None), + ) + + self.assertEqual(response.code, 200) + self.assertEqual(response.data["strategy"], "qa_pair") + + async def test_confirm_endpoint_returns_ingestion_result(self): + with tempfile.TemporaryDirectory() as tmp: + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=AsyncMock(), + embedder=AsyncMock(return_value=[[0.0] * 768]), + publisher=AsyncMock(), + ) + preview = await service.preview("faq.md", b"Q: Q\nA: A", strategy="qa_pair") + request = FakeRequest( + service, + json_data={ + "upload_id": preview["upload_id"], + "title": "FAQ", + "doc_id": "doc-1", + "collection_name": "fin_faq", + "strategy": "qa_pair", + }, + ) + + response = await confirm_document_upload( + request, + SimpleNamespace(user_type="ADMIN", employee_role=None), + ) + + self.assertEqual(response.code, 200) + self.assertEqual(response.data["doc_id"], "doc-1") + + async def test_delete_endpoint_returns_deleted_document(self): + with tempfile.TemporaryDirectory() as tmp: + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=AsyncMock(), + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + request = FakeRequest(service, json_data={"doc_id": "doc-delete-1"}) + + response = await delete_document( + request, + "doc-delete-1", + SimpleNamespace(user_type="ADMIN", employee_role=None), + ) + + self.assertEqual(response.code, 200) + self.assertEqual(response.data, {"doc_id": "doc-delete-1", "deleted": True}) + + async def test_list_endpoint_returns_documents(self): + with tempfile.TemporaryDirectory() as tmp: + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=AsyncMock(), + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + service.list_documents = AsyncMock(return_value=[{"doc_id": "doc-1"}]) + response = await list_documents( + FakeRequest(service), + SimpleNamespace(user_type="ADMIN", employee_role=None), + ) + + self.assertEqual(response.code, 200) + self.assertEqual(response.data, [{"doc_id": "doc-1"}]) + + async def test_detail_endpoint_returns_document(self): + with tempfile.TemporaryDirectory() as tmp: + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=AsyncMock(), + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + service.get_document = AsyncMock(return_value={"doc_id": "doc-1"}) + response = await get_document( + FakeRequest(service), + "doc-1", + SimpleNamespace(user_type="ADMIN", employee_role=None), + ) + + self.assertEqual(response.code, 200) + self.assertEqual(response.data, {"doc_id": "doc-1"}) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_knowledge_upload.py b/tests/test_knowledge_upload.py new file mode 100644 index 0000000..acb6980 --- /dev/null +++ b/tests/test_knowledge_upload.py @@ -0,0 +1,290 @@ +import tempfile +import unittest +import os +import time +from pathlib import Path +from unittest.mock import AsyncMock + +from rag.events import KNOWLEDGE_UPDATE_EVENT +from service.knowledge_base.upload import KnowledgeUploadService, UploadValidationError + + +class KnowledgeUploadTests(unittest.IsolatedAsyncioTestCase): + async def test_preview_saves_upload_and_returns_cleaning_and_chunk_preview(self): + with tempfile.TemporaryDirectory() as tmp: + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=AsyncMock(), + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + + result = await service.preview( + filename="notice.md", + content="\ufeff# Notice\r\n\r\n正文".encode(), + strategy="chapter_semantic", + ) + + self.assertTrue(result["upload_id"]) + self.assertTrue(result["chunks"]) + self.assertTrue(Path(tmp, result["stored_filename"]).is_file()) + self.assertIn("cleaning_warnings", result) + + async def test_confirm_ingests_then_publishes_update_and_removes_temp_file(self): + with tempfile.TemporaryDirectory() as tmp: + publisher = AsyncMock() + milvus = AsyncMock() + embedder = AsyncMock(return_value=[[0.0] * 768]) + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=milvus, + embedder=embedder, + publisher=publisher, + ) + preview = await service.preview( + filename="faq.md", + content="Q: 什么是基金?\nA: 一种集合投资工具。".encode(), + strategy="qa_pair", + ) + + result = await service.confirm( + upload_id=preview["upload_id"], + title="FAQ", + doc_id="doc-upload-1", + collection_name="fin_faq", + strategy="qa_pair", + ) + + self.assertEqual(result["doc_id"], "doc-upload-1") + publisher.assert_awaited_once() + self.assertEqual(publisher.await_args.args[0], KNOWLEDGE_UPDATE_EVENT) + self.assertFalse(Path(tmp, preview["stored_filename"]).exists()) + + async def test_confirm_passes_custom_chunk_config_to_ingestion(self): + with tempfile.TemporaryDirectory() as tmp: + milvus = AsyncMock() + embedder = AsyncMock(side_effect=lambda texts: [[0.0] * 768 for _ in texts]) + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=milvus, + embedder=embedder, + publisher=AsyncMock(), + ) + preview = await service.preview( + "doc.md", "一二三四五六七八九十十一十二".encode(), strategy="default" + ) + + await service.confirm( + upload_id=preview["upload_id"], + title="Doc", + doc_id="doc-config", + collection_name="fin_fund_doc", + strategy="default", + chunk_size=10, + chunk_overlap=2, + ) + + self.assertEqual(len(embedder.await_args.args[0]), 2) + + async def test_confirm_removes_vectors_when_update_event_publish_fails(self): + with tempfile.TemporaryDirectory() as tmp: + milvus = AsyncMock() + publisher = AsyncMock(side_effect=RuntimeError("redis unavailable")) + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=milvus, + embedder=AsyncMock(return_value=[[0.0] * 768]), + publisher=publisher, + ) + preview = await service.preview( + "policy.md", b"policy text", strategy="default" + ) + + with self.assertRaises(RuntimeError): + await service.confirm( + upload_id=preview["upload_id"], + title="Policy", + doc_id="doc-event-failure", + collection_name="fin_policy", + strategy="default", + ) + + self.assertEqual(milvus.delete.await_count, 3) + + async def test_rejects_unsupported_extension_and_oversized_file(self): + with tempfile.TemporaryDirectory() as tmp: + service = KnowledgeUploadService( + storage_dir=tmp, + max_upload_bytes=4, + milvus_client=AsyncMock(), + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + + with self.assertRaises(UploadValidationError): + await service.preview("file.exe", b"ok", strategy="default") + with self.assertRaises(UploadValidationError): + await service.preview("file.md", b"12345", strategy="default") + + async def test_confirm_rejects_unknown_collection_or_strategy_mismatch(self): + with tempfile.TemporaryDirectory() as tmp: + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=AsyncMock(), + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + preview = await service.preview( + "faq.md", b"Q: Q\nA: A", strategy="qa_pair" + ) + + with self.assertRaises(UploadValidationError): + await service.confirm( + upload_id=preview["upload_id"], title="FAQ", doc_id="d1", + collection_name="evil_collection", strategy="qa_pair", + ) + + async def test_confirm_rejects_duplicate_doc_id_before_embedding(self): + with tempfile.TemporaryDirectory() as tmp: + embedder = AsyncMock() + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=AsyncMock(), + embedder=embedder, + publisher=AsyncMock(), + document_exists=lambda doc_id: True, + ) + preview = await service.preview( + "faq.md", b"Q: Q\nA: A", strategy="qa_pair" + ) + + with self.assertRaises(UploadValidationError): + await service.confirm( + upload_id=preview["upload_id"], title="FAQ", doc_id="d1", + collection_name="fin_faq", strategy="qa_pair", + ) + embedder.assert_not_awaited() + + async def test_preview_storage_cleanup_removes_expired_upload_and_manifest(self): + with tempfile.TemporaryDirectory() as tmp: + service = KnowledgeUploadService( + storage_dir=tmp, + upload_ttl_seconds=60, + milvus_client=AsyncMock(), + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + preview = await service.preview("faq.md", b"Q: Q\nA: A", strategy="qa_pair") + path = Path(tmp, preview["stored_filename"]) + manifest = Path(tmp, f"{preview['upload_id']}.json") + old = time.time() - 120 + os.utime(path, (old, old)) + os.utime(manifest, (old, old)) + + removed = service.cleanup_expired_uploads(now=time.time()) + + self.assertEqual(removed, 1) + self.assertFalse(path.exists()) + self.assertFalse(manifest.exists()) + + async def test_confirm_rejects_expired_upload(self): + with tempfile.TemporaryDirectory() as tmp: + service = KnowledgeUploadService( + storage_dir=tmp, + upload_ttl_seconds=60, + milvus_client=AsyncMock(), + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + preview = await service.preview("faq.md", b"Q: Q\nA: A", strategy="qa_pair") + path = Path(tmp, preview["stored_filename"]) + old = time.time() - 120 + os.utime(path, (old, old)) + os.utime(Path(tmp, f"{preview['upload_id']}.json"), (old, old)) + + with self.assertRaises(UploadValidationError): + await service.confirm( + upload_id=preview["upload_id"], + title="FAQ", + doc_id="expired-doc", + collection_name="fin_faq", + strategy="qa_pair", + ) + + async def test_delete_document_removes_vectors_and_publishes_update(self): + with tempfile.TemporaryDirectory() as tmp: + milvus = AsyncMock() + publisher = AsyncMock() + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=milvus, + embedder=AsyncMock(), + publisher=publisher, + ) + + result = await service.delete_document("doc-delete-1") + + self.assertEqual(result, {"doc_id": "doc-delete-1", "deleted": True}) + self.assertEqual(milvus.delete.await_count, 3) + publisher.assert_awaited_once() + self.assertEqual(publisher.await_args.args[0], KNOWLEDGE_UPDATE_EVENT) + self.assertEqual(publisher.await_args.args[1]["doc_id"], "doc-delete-1") + self.assertEqual(publisher.await_args.args[1]["chunk_count"], 0) + self.assertEqual(publisher.await_args.args[1]["action"], "deleted") + + async def test_list_documents_groups_chunks_across_collections(self): + with tempfile.TemporaryDirectory() as tmp: + milvus = AsyncMock() + + async def query(**kwargs): + if kwargs["collection_name"] == "fin_faq": + return [ + { + "doc_id": "doc-1", + "title": "FAQ", + "section_title": "基金基础", + "strategy": "qa_pair", + }, + { + "doc_id": "doc-1", + "title": "FAQ", + "section_title": "购买流程", + "strategy": "qa_pair", + }, + ] + return [] + + milvus.query.side_effect = query + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=milvus, + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + + result = await service.list_documents() + + self.assertEqual(result, [{ + "doc_id": "doc-1", + "title": "FAQ", + "collection_name": "fin_faq", + "strategy": "qa_pair", + "chunk_count": 2, + }]) + + async def test_get_document_returns_not_found_when_doc_id_is_missing(self): + with tempfile.TemporaryDirectory() as tmp: + milvus = AsyncMock() + milvus.query.return_value = [] + service = KnowledgeUploadService( + storage_dir=tmp, + milvus_client=milvus, + embedder=AsyncMock(), + publisher=AsyncMock(), + ) + + self.assertIsNone(await service.get_document("missing-doc")) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_milvus_collections.py b/tests/test_milvus_collections.py new file mode 100644 index 0000000..0d68157 --- /dev/null +++ b/tests/test_milvus_collections.py @@ -0,0 +1,52 @@ +import unittest +from unittest.mock import AsyncMock, call +from unittest.mock import patch + +from rag.milvus_collections import KNOWLEDGE_COLLECTIONS, ensure_collections + + +class MilvusCollectionTests(unittest.IsolatedAsyncioTestCase): + async def test_creates_all_project_knowledge_collections(self): + client = AsyncMock() + client.has_collection.return_value = False + + await ensure_collections(client) + + self.assertEqual( + client.has_collection.await_args_list, + [call(name) for name in KNOWLEDGE_COLLECTIONS], + ) + self.assertEqual(client.create_collection.await_count, len(KNOWLEDGE_COLLECTIONS)) + + async def test_reuses_existing_collections(self): + client = AsyncMock() + client.has_collection.return_value = True + + await ensure_collections(client) + + client.create_collection.assert_not_awaited() + + async def test_collection_schema_contains_metadata_and_768_vector(self): + client = AsyncMock() + client.has_collection.return_value = False + + await ensure_collections(client) + + schema = client.create_collection.await_args.kwargs["schema"] + fields = {field["name"] for field in schema.to_dict()["fields"]} + self.assertTrue({"chunk_id", "doc_id", "title", "section_title", "text", "strategy", "vector"} <= fields) + vector = next(field for field in schema.to_dict()["fields"] if field["name"] == "vector") + self.assertEqual(vector["params"]["dim"], 768) + + async def test_default_path_uses_configured_milvus_client(self): + client = AsyncMock() + client.has_collection.return_value = True + + with patch("rag.milvus_collections.configured_milvus_client", return_value=client): + await ensure_collections() + + self.assertEqual(client.has_collection.await_count, len(KNOWLEDGE_COLLECTIONS)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_milvus_config.py b/tests/test_milvus_config.py new file mode 100644 index 0000000..1f4346c --- /dev/null +++ b/tests/test_milvus_config.py @@ -0,0 +1,34 @@ +import unittest +from unittest.mock import AsyncMock + +from config.settings import settings +from config.database.milvus import ensure_database + + +class MilvusConfigTests(unittest.TestCase): + def test_uses_project_database_from_milvus_db_env(self): + self.assertEqual(settings.milvus.db_name, "mutual_fund") + + +class MilvusDatabaseTests(unittest.IsolatedAsyncioTestCase): + async def test_creates_project_database_only_when_missing(self): + client = AsyncMock() + client.list_databases.return_value = ["default"] + + await ensure_database(client) + + client.create_database.assert_awaited_once_with("mutual_fund") + client.drop_database.assert_not_called() + + async def test_reuses_existing_project_database_without_recreating_it(self): + client = AsyncMock() + client.list_databases.return_value = ["default", "mutual_fund"] + + await ensure_database(client) + + client.create_database.assert_not_called() + client.drop_database.assert_not_called() + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_mock_ingest.py b/tests/test_mock_ingest.py new file mode 100644 index 0000000..aaa336b --- /dev/null +++ b/tests/test_mock_ingest.py @@ -0,0 +1,52 @@ +import unittest +from pathlib import Path +from unittest.mock import AsyncMock + +from rag.mock_ingest import MOCK_DOCUMENTS, ingest_mock_documents +from rag.milvus_delete import delete_document_vectors + + +class MockIngestTests(unittest.IsolatedAsyncioTestCase): + async def test_mock_documents_are_chunked_embedded_and_inserted(self): + milvus = AsyncMock() + embedder = AsyncMock() + embedder.return_value = [[0.0] * 768] + + await ingest_mock_documents(milvus, embedder=embedder) + + self.assertEqual(len(MOCK_DOCUMENTS), 3) + self.assertEqual(embedder.await_count, len(MOCK_DOCUMENTS)) + self.assertEqual(milvus.insert.await_count, len(MOCK_DOCUMENTS)) + inserted = [call.kwargs for call in milvus.insert.await_args_list] + self.assertEqual( + {item["collection_name"] for item in inserted}, + {item.collection for item in MOCK_DOCUMENTS}, + ) + for call_args in inserted: + rows = call_args["data"] + self.assertTrue(rows) + self.assertTrue({"doc_id", "title", "section_title", "text", "strategy", "vector"} <= rows[0].keys()) + self.assertEqual(len(rows[0]["vector"]), 768) + + def test_mock_markdown_files_are_present(self): + for item in MOCK_DOCUMENTS: + self.assertTrue(Path(item.path).is_file()) + self.assertEqual(Path(item.path).suffix, ".md") + faq = next(item for item in MOCK_DOCUMENTS if item.collection == "fin_faq") + self.assertEqual( + sum(line.startswith("Q:") for line in Path(faq.path).read_text(encoding="utf-8").splitlines()), + 40, + ) + + async def test_deletes_document_rows_by_doc_id_from_all_knowledge_collections(self): + milvus = AsyncMock() + + await delete_document_vectors("doc-123", milvus_client=milvus) + + self.assertEqual(milvus.delete.await_count, 3) + for call_args in milvus.delete.await_args_list: + self.assertEqual(call_args.kwargs["filter"], 'doc_id == "doc-123"') + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_preview.py b/tests/test_preview.py new file mode 100644 index 0000000..d51753b --- /dev/null +++ b/tests/test_preview.py @@ -0,0 +1,54 @@ +import tempfile +import unittest +from pathlib import Path + +from rag.preview import preview_document + + +class PreviewTests(unittest.TestCase): + def test_previews_manually_selected_strategy_with_custom_config(self): + with tempfile.TemporaryDirectory() as tmp: + path = Path(tmp) / "faq.md" + path.write_text("Q: 什么是净值?\nA: 基金单位价值。", encoding="utf-8") + + preview = preview_document(path, strategy="qa_pair", chunk_size=64, chunk_overlap=4) + + self.assertEqual(preview["strategy"], "qa_pair") + self.assertEqual(preview["actual_strategy"], "qa_pair") + self.assertEqual(len(preview["chunks"]), 1) + self.assertEqual(preview["chunks"][0]["section_title"], None) + + def test_rejects_auto_strategy(self): + with tempfile.TemporaryDirectory() as tmp: + path = Path(tmp) / "faq.md" + path.write_text("普通文本", encoding="utf-8") + + with self.assertRaises(ValueError): + preview_document(path, strategy="auto") + + def test_preview_reports_chapter_fallback(self): + with tempfile.TemporaryDirectory() as tmp: + path = Path(tmp) / "notice.md" + path.write_text("没有标题的普通文本", encoding="utf-8") + + preview = preview_document(path, strategy="chapter_semantic") + + self.assertEqual(preview["actual_strategy"], "default") + self.assertTrue(preview["degraded"]) + self.assertTrue(preview["warning"]) + + def test_preview_exposes_cleaning_warnings_and_uses_cleaned_text(self): + with tempfile.TemporaryDirectory() as tmp: + path = Path(tmp) / "notice.md" + path.write_text("\ufeff第一段\r\n\r\n\r\n第二段\t内容", encoding="utf-8") + + preview = preview_document(path, strategy="default") + + self.assertTrue(preview["cleaning_changed"]) + self.assertTrue(preview["cleaning_warnings"]) + self.assertEqual(preview["chunks"][0]["text"], "第一段") + self.assertEqual(preview["chunks"][1]["text"], "第二段 内容") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_retrieve.py b/tests/test_retrieve.py new file mode 100644 index 0000000..642043a --- /dev/null +++ b/tests/test_retrieve.py @@ -0,0 +1,102 @@ +import unittest +from unittest.mock import AsyncMock + +from rag.retrieve import rag_retrieve, retrieve_candidates, retrieve_with_status + + +class RetrievalTests(unittest.IsolatedAsyncioTestCase): + async def test_reads_topk_threshold_and_searches_three_business_collections(self): + milvus = AsyncMock() + milvus.search.return_value = [ + [{"id": "chunk-1", "distance": 0.9, "entity": {"doc_id": "d1"}}] + ] + embedder = AsyncMock(return_value=[[0.0] * 768]) + values = { + "agent.customer.rag.topk.faq": "3", + "agent.customer.rag.threshold.faq": "0.75", + "agent.customer.rag.topk.funddoc": "5", + "agent.customer.rag.threshold.funddoc": "0.7", + "agent.customer.rag.topk.policy": "5", + "agent.customer.rag.threshold.policy": "0.7", + } + + await retrieve_candidates( + "基金是什么", None, milvus_client=milvus, embedder=embedder, + config_getter=values.get, + ) + + self.assertEqual(milvus.search.await_count, 3) + calls = {call.kwargs["collection_name"]: call.kwargs for call in milvus.search.await_args_list} + self.assertEqual(calls["fin_faq"]["limit"], 3) + self.assertEqual(calls["fin_faq"]["filter"], "") + self.assertEqual(calls["fin_faq"]["data"], [[0.0] * 768]) + + async def test_customer_memory_is_searched_only_for_a_customer(self): + milvus = AsyncMock() + milvus.search.return_value = [[]] + embedder = AsyncMock(return_value=[[0.0] * 768]) + values = {"agent.customer.rag.topk.memory": "5", "agent.customer.rag.threshold.memory": "0.6"} + + await retrieve_candidates( + "风险", "customer-7", milvus_client=milvus, embedder=embedder, + config_getter=values.get, + ) + + memory_call = milvus.search.await_args_list[-1].kwargs + self.assertEqual(memory_call["collection_name"], "customer_memory") + self.assertEqual(memory_call["filter"], 'customer_id == "customer-7"') + + async def test_public_retrieve_returns_milvus_sources_and_filters_low_scores(self): + milvus = AsyncMock() + milvus.search.side_effect = [ + [[ + {"id": "c1", "distance": 0.90, "entity": { + "doc_id": "doc-1", "title": "FAQ", "section_title": "", + "text": "基金正文", + }}, + {"id": "c2", "distance": 0.50, "entity": { + "doc_id": "doc-low", "title": "低分", "text": "不应返回", + }}, + ]], + [[]], + [[]], + ] + embedder = AsyncMock(return_value=[[0.0] * 768]) + + sources = await rag_retrieve( + "基金", None, milvus_client=milvus, embedder=embedder, + config_getter={}.get, + ) + + self.assertEqual(sources, [{ + "doc_id": "doc-1", "title": "FAQ", "section_title": None, + "chunk_text": "基金正文", "score": 0.90, + }]) + + async def test_milvus_failure_returns_empty_sources_without_fake_hit(self): + milvus = AsyncMock() + milvus.search.side_effect = TimeoutError("Milvus timeout") + embedder = AsyncMock(return_value=[[0.0] * 768]) + + result = await retrieve_with_status( + "基金", None, milvus_client=milvus, embedder=embedder, + config_getter={}.get, + ) + + self.assertEqual(result, {"status": "milvus_unavailable", "sources": []}) + + async def test_embedding_failure_returns_failed_status_without_mysql_fallback(self): + milvus = AsyncMock() + embedder = AsyncMock(side_effect=RuntimeError("embedding down")) + + result = await retrieve_with_status( + "基金", None, milvus_client=milvus, embedder=embedder, + config_getter={}.get, + ) + + self.assertEqual(result, {"status": "embedding_failed", "sources": []}) + milvus.search.assert_not_awaited() + + +if __name__ == "__main__": + unittest.main() -- 2.54.0 From 78e326db59daf0fdf2a80e4d2a4d98370b49a5ad Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9D=92=E5=8F=B2=E8=8D=92=E4=B8=98?= Date: Fri, 11 Sep 2026 11:05:52 +0800 Subject: [PATCH 3/3] =?UTF-8?q?feat:=E6=96=B0=E5=A2=9E=E6=8E=A5=E5=8F=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- api/routers/account.py | 31 +++++++++ api/routers/holdings.py | 20 ++++++ api/routers/purchase.py | 23 +++++++ api/routers/redeem.py | 22 +++++++ model/fin_account.py | 33 ++++++++++ model/fin_holdings.py | 29 ++++++++ repositories/fin_account.py | 97 +++++++++++++++++++++++++++ repositories/fin_customer_profile.py | 21 ++++++ repositories/fin_holdings.py | 75 +++++++++++++++++++++ repositories/fin_product.py | 7 ++ schemas/account.py | 34 ++++++++++ schemas/holdings.py | 30 +++++++++ schemas/purchase.py | 24 +++++++ schemas/redeem.py | 24 +++++++ service/account.py | 80 ++++++++++++++++++++++ service/holdings.py | 37 +++++++++++ service/purchase.py | 99 ++++++++++++++++++++++++++++ service/redeem.py | 80 ++++++++++++++++++++++ 18 files changed, 766 insertions(+) create mode 100644 api/routers/account.py create mode 100644 api/routers/holdings.py create mode 100644 api/routers/purchase.py create mode 100644 api/routers/redeem.py create mode 100644 model/fin_account.py create mode 100644 model/fin_holdings.py create mode 100644 repositories/fin_account.py create mode 100644 repositories/fin_customer_profile.py create mode 100644 repositories/fin_holdings.py create mode 100644 repositories/fin_product.py create mode 100644 schemas/account.py create mode 100644 schemas/holdings.py create mode 100644 schemas/purchase.py create mode 100644 schemas/redeem.py create mode 100644 service/account.py create mode 100644 service/holdings.py create mode 100644 service/purchase.py create mode 100644 service/redeem.py diff --git a/api/routers/account.py b/api/routers/account.py new file mode 100644 index 0000000..cb1f628 --- /dev/null +++ b/api/routers/account.py @@ -0,0 +1,31 @@ +"""资金账户路由:余额查询、余额加减(业务在 service/account.py,路由只做编排)。""" +from fastapi import APIRouter, Depends +from sqlalchemy.ext.asyncio import AsyncSession + +from api.deps import get_current_user +from config.deps import get_db +from model.sys_user import SysUser +from schemas.account import AdjustReq +from service import account as account_service +from utils.response import success + +router = APIRouter() + + +@router.get("/account/balance", summary="查询当前用户余额") +async def balance( + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + result = await account_service.get_balance(db, user) + return success(result.model_dump(mode="json")) + + +@router.post("/account/adjust", summary="调整余额(add=充值 / sub=提现)") +async def adjust( + req: AdjustReq, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + result = await account_service.adjust_balance(db, user, req.direction, req.amount) + return success(result.model_dump(mode="json")) diff --git a/api/routers/holdings.py b/api/routers/holdings.py new file mode 100644 index 0000000..2ae4ac9 --- /dev/null +++ b/api/routers/holdings.py @@ -0,0 +1,20 @@ +"""持仓路由:查询当前用户持仓(业务在 service/holdings.py,路由只做编排)。""" +from fastapi import APIRouter, Depends +from sqlalchemy.ext.asyncio import AsyncSession + +from api.deps import get_current_user +from config.deps import get_db +from model.sys_user import SysUser +from service import holdings as holdings_service +from utils.response import success + +router = APIRouter() + + +@router.get("/holdings", summary="查询当前用户持仓(持有中)") +async def list_holdings( + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + result = await holdings_service.get_holdings(db, user) + return success([h.model_dump(mode="json") for h in result]) diff --git a/api/routers/purchase.py b/api/routers/purchase.py new file mode 100644 index 0000000..944155b --- /dev/null +++ b/api/routers/purchase.py @@ -0,0 +1,23 @@ + +"""申购路由:申购基金产品(业务在 service/purchase.py,路由只做编排)。""" +from fastapi import APIRouter, Depends +from sqlalchemy.ext.asyncio import AsyncSession + +from api.deps import get_current_user +from config.deps import get_db +from model.sys_user import SysUser +from schemas.purchase import PurchaseReq +from service import purchase as purchase_service +from utils.response import success + +router = APIRouter() + + +@router.post("/purchase", summary="申购基金产品") +async def purchase( + req: PurchaseReq, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + result = await purchase_service.purchase(db, user, req.product_id, req.amount) + return success(result.model_dump(mode="json")) diff --git a/api/routers/redeem.py b/api/routers/redeem.py new file mode 100644 index 0000000..768ce79 --- /dev/null +++ b/api/routers/redeem.py @@ -0,0 +1,22 @@ +"""赎回路由:赎回基金产品(业务在 service/redeem.py,路由只做编排)。""" +from fastapi import APIRouter, Depends +from sqlalchemy.ext.asyncio import AsyncSession + +from api.deps import get_current_user +from config.deps import get_db +from model.sys_user import SysUser +from schemas.redeem import RedeemReq +from service import redeem as redeem_service +from utils.response import success + +router = APIRouter() + + +@router.post("/redeem", summary="赎回基金产品") +async def redeem( + req: RedeemReq, + user: SysUser = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +): + result = await redeem_service.redeem(db, user, req.product_id, req.shares) + return success(result.model_dump(mode="json")) diff --git a/model/fin_account.py b/model/fin_account.py new file mode 100644 index 0000000..207d923 --- /dev/null +++ b/model/fin_account.py @@ -0,0 +1,33 @@ +"""fin_account 客户资金账户表 ORM 模型(现金余额,一人一户)。 + +balance 为账户总余额(含冻结部分),可用余额 = balance - frozen_amount,不落库。 +""" +from __future__ import annotations + +from datetime import datetime +from decimal import Decimal + +from sqlalchemy import BigInteger, DateTime, Integer, Numeric, String, func +from sqlalchemy.orm import Mapped, mapped_column + +from model.base import Base + + +class FinAccount(Base): + __tablename__ = "fin_account" + __table_args__ = {"comment": "客户资金账户表(现金余额,申购扣款/赎回入账的账务载体)"} + + customer_id: Mapped[int] = mapped_column( + BigInteger, primary_key=True, autoincrement=False + ) + balance: Mapped[Decimal] = mapped_column(Numeric(18, 2), server_default="0") + frozen_amount: Mapped[Decimal] = mapped_column(Numeric(18, 2), server_default="0") + currency: Mapped[str] = mapped_column(String(8), server_default="CNY") + status: Mapped[str] = mapped_column(String(16), server_default="正常") + version: Mapped[int] = mapped_column(Integer, server_default="0") + create_time: Mapped[datetime] = mapped_column( + DateTime, server_default=func.now() + ) + update_time: Mapped[datetime] = mapped_column( + DateTime, server_default=func.now(), onupdate=func.now() + ) diff --git a/model/fin_holdings.py b/model/fin_holdings.py new file mode 100644 index 0000000..b312c0e --- /dev/null +++ b/model/fin_holdings.py @@ -0,0 +1,29 @@ +"""fin_holdings 持仓表 ORM 模型(当前/历史持仓快照)。""" +from __future__ import annotations + +from datetime import datetime +from decimal import Decimal + +from sqlalchemy import BigInteger, DateTime, Numeric, String, func +from sqlalchemy.orm import Mapped, mapped_column + +from model.base import Base + + +class FinHoldings(Base): + __tablename__ = "fin_holdings" + __table_args__ = {"comment": "持仓表(当前/历史持仓快照)"} + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + customer_id: Mapped[int] = mapped_column(BigInteger) + product_id: Mapped[int] = mapped_column(BigInteger) + shares: Mapped[Decimal] = mapped_column(Numeric(18, 4), server_default="0") + cost_amount: Mapped[Decimal] = mapped_column(Numeric(18, 2), server_default="0") + current_value: Mapped[Decimal] = mapped_column(Numeric(18, 2), server_default="0") + profit_loss: Mapped[Decimal] = mapped_column(Numeric(18, 2), server_default="0") + profit_ratio: Mapped[Decimal] = mapped_column(Numeric(8, 4), server_default="0") + status: Mapped[str] = mapped_column(String(16), server_default="持有中") + create_time: Mapped[datetime] = mapped_column(DateTime, server_default=func.now()) + update_time: Mapped[datetime] = mapped_column( + DateTime, server_default=func.now(), onupdate=func.now() + ) diff --git a/repositories/fin_account.py b/repositories/fin_account.py new file mode 100644 index 0000000..831d943 --- /dev/null +++ b/repositories/fin_account.py @@ -0,0 +1,97 @@ +"""fin_account 仓储:按客户 ID 取资金账户 + 余额加减(原子 UPDATE)。 + +注:本表主键为 customer_id(非 id),故不复用 BaseRepository.delete/count 中的 id 约定。 +余额增减用原子 UPDATE(balance = balance ± delta)防并发丢更新,不依赖乐观锁重试。 +""" +from __future__ import annotations + +from decimal import Decimal + +from sqlalchemy import select, update + +from model.fin_account import FinAccount +from repositories.base import BaseRepository + + +class FinAccountRepo(BaseRepository): + model = FinAccount + + async def get_by_customer_id(self, customer_id: int) -> FinAccount | None: + return await self.db.scalar( + select(FinAccount).where(FinAccount.customer_id == customer_id) + ) + + async def create(self, customer_id: int, balance: Decimal) -> FinAccount: + """开资金户(充值时账户不存在则自动开户入账)。""" + account = FinAccount(customer_id=customer_id, balance=balance) + self.db.add(account) + await self.db.commit() + await self.db.refresh(account) + return account + + async def add_balance(self, customer_id: int, delta: Decimal) -> FinAccount: + """入账:balance += delta,原子自增后回读最新余额。""" + await self.db.execute( + update(FinAccount) + .where(FinAccount.customer_id == customer_id) + .values( + balance=FinAccount.balance + delta, + version=FinAccount.version + 1, + ) + ) + await self.db.commit() + return await self.get_by_customer_id(customer_id) + + async def subtract_balance( + self, customer_id: int, delta: Decimal + ) -> FinAccount | None: + """出账:balance -= delta,可用余额(balance - frozen_amount)不足时返回 None。""" + result = await self.db.execute( + update(FinAccount) + .where( + FinAccount.customer_id == customer_id, + FinAccount.balance - FinAccount.frozen_amount >= delta, + ) + .values( + balance=FinAccount.balance - delta, + version=FinAccount.version + 1, + ) + ) + await self.db.commit() + if result.rowcount == 0: + return None + return await self.get_by_customer_id(customer_id) + + async def deduct_balance(self, customer_id: int, delta: Decimal) -> bool: + """申购事务内扣款:balance -= delta(可用余额不足则不动)。 + + 不 commit,由 service 层事务统一提交,保证「扣款 + 加仓」原子性。 + 可用余额(balance - frozen_amount)不足时返回 False。 + """ + result = await self.db.execute( + update(FinAccount) + .where( + FinAccount.customer_id == customer_id, + FinAccount.balance - FinAccount.frozen_amount >= delta, + ) + .values( + balance=FinAccount.balance - delta, + version=FinAccount.version + 1, + ) + ) + return result.rowcount > 0 + + async def credit_balance(self, customer_id: int, delta: Decimal) -> bool: + """赎回事务内入账:balance += delta。 + + 不 commit,由 service 层事务统一提交,保证「减仓 + 入账」原子性。 + """ + result = await self.db.execute( + update(FinAccount) + .where(FinAccount.customer_id == customer_id) + .values( + balance=FinAccount.balance + delta, + version=FinAccount.version + 1, + ) + ) + return result.rowcount > 0 diff --git a/repositories/fin_customer_profile.py b/repositories/fin_customer_profile.py new file mode 100644 index 0000000..d168254 --- /dev/null +++ b/repositories/fin_customer_profile.py @@ -0,0 +1,21 @@ +"""fin_customer_profile 画像仓储:按客户 ID 取画像。 + +注:本表主键为 customer_id(非 id),故不复用 BaseRepository.get 的 id 约定。 +""" +from __future__ import annotations + +from sqlalchemy import select + +from model.fin_customer_profile import FinCustomerProfile +from repositories.base import BaseRepository + + +class FinCustomerProfileRepo(BaseRepository): + model = FinCustomerProfile + + async def get_by_customer_id(self, customer_id: int) -> FinCustomerProfile | None: + return await self.db.scalar( + select(FinCustomerProfile).where( + FinCustomerProfile.customer_id == customer_id + ) + ) diff --git a/repositories/fin_holdings.py b/repositories/fin_holdings.py new file mode 100644 index 0000000..253a2a9 --- /dev/null +++ b/repositories/fin_holdings.py @@ -0,0 +1,75 @@ +"""fin_holdings 持仓仓储:按客户(+状态)查持仓、申购加仓 upsert、赎回减仓。""" +from __future__ import annotations + +from decimal import Decimal + +from sqlalchemy import case, select, update +from sqlalchemy.dialects.mysql import insert as mysql_insert + +from model.fin_holdings import FinHoldings +from repositories.base import BaseRepository + + +class FinHoldingsRepo(BaseRepository): + model = FinHoldings + + async def list_by_customer( + self, customer_id: int, status: str | None = None + ) -> list[FinHoldings]: + """按客户 ID 查持仓,可按状态过滤(status=None 表示不过滤)。""" + stmt = select(FinHoldings).where(FinHoldings.customer_id == customer_id) + if status is not None: + stmt = stmt.where(FinHoldings.status == status) + stmt = stmt.order_by(FinHoldings.id) + return list((await self.db.scalars(stmt)).all()) + + async def get_by_customer_product( + self, customer_id: int, product_id: int + ) -> FinHoldings | None: + return await self.db.scalar( + select(FinHoldings).where( + FinHoldings.customer_id == customer_id, + FinHoldings.product_id == product_id, + ) + ) + + async def upsert( + self, customer_id: int, product_id: int, add_shares: Decimal, add_cost: Decimal + ) -> None: + """申购加仓:有则加份额/成本,无则新增(靠 uk_customer_product 唯一键)。不 commit。""" + stmt = mysql_insert(FinHoldings).values( + customer_id=customer_id, + product_id=product_id, + shares=add_shares, + cost_amount=add_cost, + ) + stmt = stmt.on_duplicate_key_update( + shares=FinHoldings.shares + add_shares, + cost_amount=FinHoldings.cost_amount + add_cost, + ) + await self.db.execute(stmt) + + async def redeem( + self, customer_id: int, product_id: int, redeem_shares: Decimal + ) -> bool: + """赎回减仓:shares -= redeem_shares,份额归 0 时状态置'已清仓'。 + + 持仓份额不足(含无持仓、shares=0)时不动作,返回 False。 + 不 commit,由 service 层事务统一提交,保证「减仓 + 入账」原子性。 + """ + result = await self.db.execute( + update(FinHoldings) + .where( + FinHoldings.customer_id == customer_id, + FinHoldings.product_id == product_id, + FinHoldings.shares >= redeem_shares, + ) + .values( + shares=FinHoldings.shares - redeem_shares, + status=case( + (FinHoldings.shares - redeem_shares == 0, "已清仓"), + else_=FinHoldings.status, + ), + ) + ) + return result.rowcount > 0 diff --git a/repositories/fin_product.py b/repositories/fin_product.py new file mode 100644 index 0000000..1b17e52 --- /dev/null +++ b/repositories/fin_product.py @@ -0,0 +1,7 @@ +"""fin_product 产品仓储:申购按主键取产品(复用 BaseRepository.get)。""" +from model.fin_product import FinProduct +from repositories.base import BaseRepository + + +class FinProductRepo(BaseRepository): + model = FinProduct diff --git a/schemas/account.py b/schemas/account.py new file mode 100644 index 0000000..cfaa2b1 --- /dev/null +++ b/schemas/account.py @@ -0,0 +1,34 @@ +"""资金账户相关 DTO。 + +金额统一序列化为两位小数字符串(如 "48000.00"),避免 JSON number 的浮点精度隐患。 +""" +from datetime import datetime +from decimal import Decimal +from typing import Literal + +from pydantic import BaseModel, Field, field_serializer + + +class BalanceResp(BaseModel): + """用户余额响应体。available_balance = balance - frozen_amount,由 service 派生。""" + + customer_id: int + balance: Decimal + available_balance: Decimal + frozen_amount: Decimal + currency: str + status: str + update_time: datetime | None = None + + @field_serializer("balance", "available_balance", "frozen_amount") + def _fmt_money(self, value: Decimal) -> str: + return f"{value:.2f}" + + +class AdjustReq(BaseModel): + """加减余额入参。amount 为正数(元),direction 决定加/减。""" + + direction: Literal["add", "sub"] = Field( + ..., description="add=入账(充值),sub=出账(提现)" + ) + amount: Decimal = Field(gt=0) diff --git a/schemas/holdings.py b/schemas/holdings.py new file mode 100644 index 0000000..437dd70 --- /dev/null +++ b/schemas/holdings.py @@ -0,0 +1,30 @@ +"""持仓相关 DTO。 + +金额字段序列化为字符串,避免 JSON number 浮点精度隐患; +份额 / 盈亏比例保留 4 位,金额保留 2 位。 +""" +from decimal import Decimal + +from pydantic import BaseModel, field_serializer + + +class HoldingResp(BaseModel): + """单条持仓记录。""" + + id: int + customer_id: int + product_id: int + shares: Decimal + cost_amount: Decimal + current_value: Decimal + profit_loss: Decimal + profit_ratio: Decimal + status: str + + @field_serializer("shares", "profit_ratio") + def _fmt_ratio(self, value: Decimal) -> str: + return f"{value:.4f}" + + @field_serializer("cost_amount", "current_value", "profit_loss") + def _fmt_money(self, value: Decimal) -> str: + return f"{value:.2f}" diff --git a/schemas/purchase.py b/schemas/purchase.py new file mode 100644 index 0000000..2cd5680 --- /dev/null +++ b/schemas/purchase.py @@ -0,0 +1,24 @@ +"""申购相关 DTO。""" +from decimal import Decimal + +from pydantic import BaseModel, Field, field_serializer + +from schemas.holdings import HoldingResp + + +class PurchaseReq(BaseModel): + """申购入参。amount 为正数(元),按净值折算份额。""" + + product_id: int + amount: Decimal = Field(gt=0) + + +class PurchaseResp(BaseModel): + """申购结果:最新余额 + 该产品最新持仓。""" + + balance: Decimal + holding: HoldingResp + + @field_serializer("balance") + def _fmt_balance(self, value: Decimal) -> str: + return f"{value:.2f}" diff --git a/schemas/redeem.py b/schemas/redeem.py new file mode 100644 index 0000000..b80071c --- /dev/null +++ b/schemas/redeem.py @@ -0,0 +1,24 @@ +"""赎回相关 DTO。""" +from decimal import Decimal + +from pydantic import BaseModel, Field, field_serializer + +from schemas.holdings import HoldingResp + + +class RedeemReq(BaseModel): + """赎回入参。shares 为正数(份额),按净值折算入账金额。""" + + product_id: int + shares: Decimal = Field(gt=0) + + +class RedeemResp(BaseModel): + """赎回结果:最新余额 + 该产品最新持仓。""" + + balance: Decimal + holding: HoldingResp + + @field_serializer("balance") + def _fmt_balance(self, value: Decimal) -> str: + return f"{value:.2f}" diff --git a/service/account.py b/service/account.py new file mode 100644 index 0000000..60d0469 --- /dev/null +++ b/service/account.py @@ -0,0 +1,80 @@ +"""资金账户服务:余额查询 + 余额加减(路由层只编排,不碰数据/逻辑)。""" +from __future__ import annotations + +from decimal import Decimal + +from sqlalchemy.ext.asyncio import AsyncSession + +from model.fin_account import FinAccount +from model.sys_user import SysUser +from repositories.fin_account import FinAccountRepo +from schemas.account import BalanceResp +from utils.exceptions import ForbiddenError, NotFoundError, ParamError + +_ZERO = Decimal("0.00") +_DEFAULT_CURRENCY = "CNY" + + +def _build_balance_resp(account: FinAccount) -> BalanceResp: + return BalanceResp( + customer_id=account.customer_id, + balance=account.balance, + available_balance=account.balance - account.frozen_amount, + frozen_amount=account.frozen_amount, + currency=account.currency, + status=account.status, + update_time=account.update_time, + ) + + +async def get_balance(db: AsyncSession, user: SysUser) -> BalanceResp: + """查询当前用户现金余额。 + + - 仅客户账号可查(员工共用 sys_user,但无资金账户); + - 未开资金户时按零余额返回,不报错。 + """ + if user.user_type != "CUSTOMER": + raise ForbiddenError("仅客户账号可查询资金余额") + + account = await FinAccountRepo(db).get_by_customer_id(user.id) + if account is None: + return BalanceResp( + customer_id=user.id, + balance=_ZERO, + available_balance=_ZERO, + frozen_amount=_ZERO, + currency=_DEFAULT_CURRENCY, + status="正常", + ) + + return _build_balance_resp(account) + + +async def adjust_balance( + db: AsyncSession, user: SysUser, direction: str, amount: Decimal +) -> BalanceResp: + """加减余额:direction=add 入账 / sub 出账。 + + - add:balance += amount,未开资金户时自动开户入账; + - sub:balance -= amount,可用余额(balance - frozen_amount)不足时报错。 + """ + if user.user_type != "CUSTOMER": + raise ForbiddenError("仅客户账号可调整余额") + + repo = FinAccountRepo(db) + account = await repo.get_by_customer_id(user.id) + + if direction == "add": + if account is None: + account = await repo.create(user.id, amount) + else: + account = await repo.add_balance(user.id, amount) + else: # sub + if account is None: + raise NotFoundError("资金账户不存在") + updated = await repo.subtract_balance(user.id, amount) + if updated is None: + raise ParamError("可用余额不足") + account = updated + + return _build_balance_resp(account) diff --git a/service/holdings.py b/service/holdings.py new file mode 100644 index 0000000..8b16926 --- /dev/null +++ b/service/holdings.py @@ -0,0 +1,37 @@ +"""持仓服务:查询当前客户持仓(路由层只编排,不碰数据/逻辑)。""" +from __future__ import annotations + +from sqlalchemy.ext.asyncio import AsyncSession + +from model.sys_user import SysUser +from repositories.fin_holdings import FinHoldingsRepo +from schemas.holdings import HoldingResp +from utils.exceptions import ForbiddenError + +_HOLDING_STATUS = "持有中" + + +async def get_holdings(db: AsyncSession, user: SysUser) -> list[HoldingResp]: + """查询当前客户的在持持仓(status=持有中)。 + + - 仅客户账号可查(员工共用 sys_user,但无持仓); + - 无持仓返回空列表,不报错。 + """ + if user.user_type != "CUSTOMER": + raise ForbiddenError("仅客户账号可查询持仓") + + holdings = await FinHoldingsRepo(db).list_by_customer(user.id, _HOLDING_STATUS) + return [ + HoldingResp( + id=h.id, + customer_id=h.customer_id, + product_id=h.product_id, + shares=h.shares, + cost_amount=h.cost_amount, + current_value=h.current_value, + profit_loss=h.profit_loss, + profit_ratio=h.profit_ratio, + status=h.status, + ) + for h in holdings + ] diff --git a/service/purchase.py b/service/purchase.py new file mode 100644 index 0000000..fa44f49 --- /dev/null +++ b/service/purchase.py @@ -0,0 +1,99 @@ +"""申购服务:风险匹配校验 + 余额扣减 + 持仓加仓(单事务原子)。""" +from __future__ import annotations + +from decimal import ROUND_HALF_UP, Decimal + +from sqlalchemy.ext.asyncio import AsyncSession + +from model.sys_user import SysUser +from repositories.fin_account import FinAccountRepo +from repositories.fin_customer_profile import FinCustomerProfileRepo +from repositories.fin_holdings import FinHoldingsRepo +from repositories.fin_product import FinProductRepo +from schemas.holdings import HoldingResp +from schemas.purchase import PurchaseResp +from utils.exceptions import ( + ForbiddenError, + NotFoundError, + NotSuitableError, + ParamError, +) + +_MONEY = Decimal("0.01") +_SHARES = Decimal("0.0001") + +# 风险等级 → 序号。兼容两套口径:R1~R5 与 保守~激进(同一映射)。 +_RISK_RANK = { + "R1": 1, "R2": 2, "R3": 3, "R4": 4, "R5": 5, + "保守": 1, "稳健": 2, "平衡": 3, "进取": 4, "激进": 5, +} + + +def _risk_rank(level: str | None) -> int | None: + """客户画像 / 产品的风险等级统一转序号;未知返回 None。""" + if not level: + return None + return _RISK_RANK.get(level.strip()) + + +async def purchase( + db: AsyncSession, user: SysUser, product_id: int, amount: Decimal +) -> PurchaseResp: + """申购基金:校验通过后扣减余额并加仓,全程单事务。 + + - 仅客户可申购; + - 产品须在售且净值非空; + - 客户风险等级序号 >= 产品风险等级序号,否则 1005 拦截; + - 余额不足拦截;扣款 + 加仓要么都成、要么都回滚。 + """ + if user.user_type != "CUSTOMER": + raise ForbiddenError("仅客户账号可申购") + + amount = amount.quantize(_MONEY, rounding=ROUND_HALF_UP) + + product = await FinProductRepo(db).get(product_id) + if product is None: + raise NotFoundError("产品不存在") + if product.status != "在售": + raise ParamError("产品不在售") + if product.nav is None: + raise ParamError("产品暂无净值,无法申购") + + profile = await FinCustomerProfileRepo(db).get_by_customer_id(user.id) + customer_rank = _risk_rank(profile.risk_level if profile else None) + if customer_rank is None: + raise NotSuitableError("客户无风险等级,无法申购") + product_rank = _risk_rank(product.risk_level) + if product_rank is None or customer_rank < product_rank: + raise NotSuitableError() + + shares = (amount / product.nav).quantize(_SHARES, rounding=ROUND_HALF_UP) + + account_repo = FinAccountRepo(db) + holdings_repo = FinHoldingsRepo(db) + + try: + if not await account_repo.deduct_balance(user.id, amount): + raise ParamError("可用余额不足") + await holdings_repo.upsert(user.id, product_id, shares, amount) + await db.commit() + except Exception: + await db.rollback() + raise + + account = await account_repo.get_by_customer_id(user.id) + holding = await holdings_repo.get_by_customer_product(user.id, product_id) + return PurchaseResp( + balance=account.balance, + holding=HoldingResp( + id=holding.id, + customer_id=holding.customer_id, + product_id=holding.product_id, + shares=holding.shares, + cost_amount=holding.cost_amount, + current_value=holding.current_value, + profit_loss=holding.profit_loss, + profit_ratio=holding.profit_ratio, + status=holding.status, + ), + ) diff --git a/service/redeem.py b/service/redeem.py new file mode 100644 index 0000000..97af761 --- /dev/null +++ b/service/redeem.py @@ -0,0 +1,80 @@ +"""赎回服务:校验持仓 → 减仓 + 余额入账(单事务原子)。""" +from __future__ import annotations + +from decimal import ROUND_HALF_UP, Decimal + +from sqlalchemy.ext.asyncio import AsyncSession + +from model.sys_user import SysUser +from repositories.fin_account import FinAccountRepo +from repositories.fin_holdings import FinHoldingsRepo +from repositories.fin_product import FinProductRepo +from schemas.holdings import HoldingResp +from schemas.redeem import RedeemResp +from utils.exceptions import ForbiddenError, NotFoundError, ParamError + +_MONEY = Decimal("0.01") +_SHARES = Decimal("0.0001") + + +async def redeem( + db: AsyncSession, user: SysUser, product_id: int, shares: Decimal +) -> RedeemResp: + """赎回基金:校验通过后减仓并按净值入账,全程单事务。 + + - 仅客户可赎回; + - 产品须存在且净值非空; + - 持仓须存在且份额 > 0,赎回份额不得超过持仓份额; + - 减仓(归 0 时置'已清仓')+ 入账要么都成、要么都回滚。 + """ + if user.user_type != "CUSTOMER": + raise ForbiddenError("仅客户账号可赎回") + + shares = shares.quantize(_SHARES, rounding=ROUND_HALF_UP) + if shares <= 0: + raise ParamError("赎回份额须大于 0") + + product = await FinProductRepo(db).get(product_id) + if product is None: + raise NotFoundError("产品不存在") + if product.nav is None: + raise ParamError("产品暂无净值,无法赎回") + + account_repo = FinAccountRepo(db) + holdings_repo = FinHoldingsRepo(db) + + holding = await holdings_repo.get_by_customer_product(user.id, product_id) + if holding is None or holding.shares <= 0: + raise ParamError("无可赎回份额") + if holding.shares < shares: + raise ParamError("可赎回份额不足") + + credited = (shares * product.nav).quantize(_MONEY, rounding=ROUND_HALF_UP) + if credited <= 0: + raise ParamError("赎回金额过低") + + try: + if not await holdings_repo.redeem(user.id, product_id, shares): + raise ParamError("可赎回份额不足") + await account_repo.credit_balance(user.id, credited) + await db.commit() + except Exception: + await db.rollback() + raise + + account = await account_repo.get_by_customer_id(user.id) + holding = await holdings_repo.get_by_customer_product(user.id, product_id) + return RedeemResp( + balance=account.balance, + holding=HoldingResp( + id=holding.id, + customer_id=holding.customer_id, + product_id=holding.product_id, + shares=holding.shares, + cost_amount=holding.cost_amount, + current_value=holding.current_value, + profit_loss=holding.profit_loss, + profit_ratio=holding.profit_ratio, + status=holding.status, + ), + ) -- 2.54.0