diff --git a/app/api/auth.py b/app/api/auth.py new file mode 100644 index 0000000..22a253e --- /dev/null +++ b/app/api/auth.py @@ -0,0 +1,25 @@ +"""开发期 Mock 登录:签发 JWT。""" + +from __future__ import annotations + +from fastapi import APIRouter, Request + +from app.gateway.jwt_service import infer_roles, issue_token +from app.model.schemas import LoginRequest, LoginResponseData +from app.utils.response import ok + +router = APIRouter(prefix="/api/auth", tags=["auth"]) + + +@router.post("/login") +def login(body: LoginRequest, request: Request): + roles = infer_roles(body.actor_id, body.token_type, body.roles) + token, expires_in = issue_token(body.actor_id, body.token_type, roles) + trace_id = getattr(request.state, "trace_id", "unknown") + data = LoginResponseData( + access_token=token, + expires_in=expires_in, + sub=body.actor_id, + roles=roles, + ) + return ok(data.model_dump(), trace_id, message="login ok") diff --git a/app/api/chat.py b/app/api/chat.py index 1243c84..6e35ee2 100644 --- a/app/api/chat.py +++ b/app/api/chat.py @@ -1 +1,70 @@ -"""对话接口:四 Agent 统一 chat 入口(按 X-Agent-Type 分流)。""" +"""对话接口:四 Agent 统一 chat 入口。""" + +from __future__ import annotations + +from typing import Annotated + +from fastapi import APIRouter, Depends, Request + +from app.gateway.auth_deps import get_auth_context +from app.gateway.ownership import assert_customer_access, resolve_effective_customer_id +from app.model.schemas import AuthContext, ChatRequest, ChatResponseData +from app.repository.agent_repository import AgentSessionRepository +from app.repository.audit_repository import AuditRepository +from app.service.agent_service import run_chat +from app.utils.input_guard import validate_user_message +from app.utils.response import ok + +router = APIRouter(prefix="/api", tags=["chat"]) + + +@router.post("/chat") +def chat( + body: ChatRequest, + request: Request, + ctx: Annotated[AuthContext, Depends(get_auth_context)], +): + message = validate_user_message(body.message) + customer_id = resolve_effective_customer_id(ctx, body.customer_id) + if customer_id and ctx.agent_type in ("advisor", "analyst", "risk", "customer"): + assert_customer_access(ctx, customer_id) + + session_repo = AgentSessionRepository() + session_id = session_repo.ensure_session(ctx, body.session_id, customer_id) + + seq = session_repo.next_seq(session_id) + session_repo.insert_message( + session_id=session_id, + trace_id=ctx.trace_id, + seq_no=seq, + role="user", + content=message, + ) + + reply, has_disclaimer = run_chat(ctx, message, customer_id) + session_repo.insert_message( + session_id=session_id, + trace_id=ctx.trace_id, + seq_no=seq + 1, + role="assistant", + content=reply, + has_disclaimer=has_disclaimer, + ) + + AuditRepository().insert( + trace_id=ctx.trace_id, + event_type="chat_completed", + agent_type=ctx.agent_type, + actor_id=ctx.sub, + customer_id=customer_id, + input_summary={"session_id": session_id, "message_len": len(message)}, + decision="success", + ) + + data = ChatResponseData( + session_id=session_id, + reply=reply, + agent_type=ctx.agent_type, + has_disclaimer=has_disclaimer, + ) + return ok(data.model_dump(), ctx.trace_id) diff --git a/app/config/database.py b/app/config/database.py index 896acc2..f57665d 100644 --- a/app/config/database.py +++ b/app/config/database.py @@ -1 +1,27 @@ -"""MySQL / Redis / Neo4j 连接工厂。""" +"""MySQL / Redis 连接工厂。""" + +from functools import lru_cache + +from sqlalchemy import create_engine +from sqlalchemy.engine import Engine + +from app.config.settings import settings + + +def _mysql_url(database: str) -> str: + pwd = settings.mysql_password + auth = f"{settings.mysql_user}:{pwd}" if pwd else settings.mysql_user + return ( + f"mysql+pymysql://{auth}@{settings.mysql_host}:{settings.mysql_port}" + f"/{database}?charset=utf8mb4" + ) + + +@lru_cache +def get_agent_engine() -> Engine: + return create_engine(_mysql_url(settings.mysql_database), pool_pre_ping=True) + + +@lru_cache +def get_core_engine() -> Engine: + return create_engine(_mysql_url(settings.mysql_core_database), pool_pre_ping=True) diff --git a/app/config/settings.py b/app/config/settings.py index b96c208..07542b9 100644 --- a/app/config/settings.py +++ b/app/config/settings.py @@ -29,5 +29,9 @@ class Settings(BaseSettings): deepseek_api_key: str = "" deepseek_base_url: str = "https://api.deepseek.com" + jwt_dev_secret: str = "change-me-in-dev-only" + jwt_dev_algorithm: str = "HS256" + jwt_dev_expire_hours: int = 8 + settings = Settings() diff --git a/app/gateway/__init__.py b/app/gateway/__init__.py new file mode 100644 index 0000000..1ca6de1 --- /dev/null +++ b/app/gateway/__init__.py @@ -0,0 +1 @@ +"""Gateway 模块。""" diff --git a/app/gateway/auth_deps.py b/app/gateway/auth_deps.py new file mode 100644 index 0000000..586c883 --- /dev/null +++ b/app/gateway/auth_deps.py @@ -0,0 +1,55 @@ +"""FastAPI 依赖:解析 JWT 并构建 AuthContext。""" + +from __future__ import annotations + +from typing import Annotated + +from fastapi import Header, Request + +from app.gateway.jwt_service import decode_token +from app.gateway.rbac import assert_agent_access +from app.model.schemas import AgentType, AuthContext +from app.utils.exceptions import UnauthorizedError + + +def _parse_bearer(authorization: str | None) -> str: + if not authorization: + raise UnauthorizedError("缺少 Authorization 头", error_code="AUTH_401_MISSING") + scheme, _, token = authorization.partition(" ") + if scheme.lower() != "bearer" or not token: + raise UnauthorizedError("Authorization 格式错误", error_code="AUTH_401_FORMAT") + return token + + +def get_auth_context( + request: Request, + authorization: Annotated[str | None, Header()] = None, + x_agent_type: Annotated[str | None, Header(alias="X-Agent-Type")] = None, +) -> AuthContext: + token = _parse_bearer(authorization) + payload = decode_token(token) + trace_id = getattr(request.state, "trace_id", None) or request.headers.get("X-Trace-Id", "unknown") + + if not x_agent_type: + raise UnauthorizedError("缺少 X-Agent-Type 头", error_code="AUTH_401_AGENT_TYPE") + + try: + agent_type: AgentType = x_agent_type # type: ignore[assignment] + if agent_type not in ("customer", "advisor", "analyst", "risk"): + raise ValueError + except ValueError as exc: + raise UnauthorizedError("X-Agent-Type 无效", error_code="AUTH_401_AGENT_TYPE") from exc + + ctx = AuthContext( + sub=str(payload["sub"]), + token_type=payload["token_type"], + roles=list(payload.get("roles") or []), + permissions=list(payload.get("permissions") or []), + tenant_id=str(payload.get("tenant_id") or "default"), + trace_id=trace_id, + agent_type=agent_type, + jti=str(payload.get("jti") or ""), + customer_id=payload.get("customer_id"), + ) + assert_agent_access(ctx) + return ctx diff --git a/app/gateway/jwt_service.py b/app/gateway/jwt_service.py new file mode 100644 index 0000000..e498b4c --- /dev/null +++ b/app/gateway/jwt_service.py @@ -0,0 +1,143 @@ +"""JWT 签发与校验(开发环境 HS256)。""" + +from __future__ import annotations + +import uuid +from datetime import UTC, datetime, timedelta +from typing import Any + +from jose import JWTError, jwt + +from app.config.settings import settings +from app.model.schemas import TokenType +from app.utils.exceptions import UnauthorizedError + +ISSUER = "https://idp.jinrong.dev" +AUDIENCE = "agent-gateway" + +ROLE_PERMISSIONS: dict[str, list[str]] = { + "customer": [ + "agent:customer:chat", + "profile:l1:read", + "profile:l1:write", + "core:holding:read:self", + ], + "advisor": [ + "agent:advisor:chat", + "profile:l1:read", + "profile:l2:read", + "profile:l2:write", + "profile:l3:read", + "core:holding:read:assigned", + ], + "analyst": [ + "agent:analyst:chat", + "profile:l1:read", + "profile:l2:read", + "profile:l3:read", + "core:holding:read:scoped", + "sql:execute:readonly", + "risk:alert:read", + ], + "risk_officer": [ + "agent:risk:chat", + "profile:l1:read", + "profile:l2:read", + "profile:l3:read", + "profile:l3:write", + "risk:alert:write", + "risk:suitability:write", + "core:holding:read:all", + ], + "compliance": ["audit:read:all", "agent:advisor:audit", "compliance:hit:read"], + "ops": ["agent:advisor:stats", "audit:read:aggregated"], + "service_risk": [ + "agent:risk:suitability_check", + "risk:suitability:write", + "profile:l1:read", + "profile:l2:read", + "audit:write", + ], +} + +DEFAULT_ROLES_BY_ACTOR: dict[str, list[str]] = { + "STAFF-10086": ["advisor"], + "STAFF-10087": ["advisor"], + "STAFF-10088": ["advisor"], + "STAFF-10089": ["advisor"], + "STAFF-10090": ["advisor"], + "STAFF-10091": ["advisor", "compliance"], + "STAFF-20001": ["analyst"], + "STAFF-20002": ["analyst"], + "STAFF-30001": ["risk_officer"], + "STAFF-30002": ["risk_officer"], + "STAFF-40001": ["compliance"], + "STAFF-40002": ["compliance"], + "STAFF-50001": ["ops"], + "CUST-9527": ["customer"], + "CUST-1001": ["customer"], + "CUST-1002": ["customer"], + "CUST-1010": ["customer"], +} + + +def infer_roles(actor_id: str, token_type: TokenType, roles: list[str] | None) -> list[str]: + if roles: + return roles + if token_type == "customer": + return ["customer"] + return DEFAULT_ROLES_BY_ACTOR.get(actor_id, ["analyst"]) + + +def permissions_for_roles(roles: list[str]) -> list[str]: + perms: list[str] = [] + seen: set[str] = set() + for role in roles: + for perm in ROLE_PERMISSIONS.get(role, []): + if perm not in seen: + seen.add(perm) + perms.append(perm) + return perms + + +def issue_token( + actor_id: str, + token_type: TokenType, + roles: list[str] | None = None, +) -> tuple[str, int]: + role_list = infer_roles(actor_id, token_type, roles) + now = datetime.now(UTC) + expires = now + timedelta(hours=settings.jwt_dev_expire_hours) + payload: dict[str, Any] = { + "iss": ISSUER, + "sub": actor_id, + "aud": AUDIENCE, + "exp": int(expires.timestamp()), + "iat": int(now.timestamp()), + "jti": str(uuid.uuid4()), + "token_type": token_type, + "roles": role_list, + "permissions": permissions_for_roles(role_list), + "tenant_id": "default", + } + if token_type == "customer": + payload["customer_id"] = actor_id + if "advisor" in role_list and token_type == "staff": + payload["advisor_id"] = actor_id + + token = jwt.encode(payload, settings.jwt_dev_secret, algorithm=settings.jwt_dev_algorithm) + return token, settings.jwt_dev_expire_hours * 3600 + + +def decode_token(token: str) -> dict[str, Any]: + try: + payload = jwt.decode( + token, + settings.jwt_dev_secret, + algorithms=[settings.jwt_dev_algorithm], + audience=AUDIENCE, + issuer=ISSUER, + ) + except JWTError as exc: + raise UnauthorizedError("Token 无效或已过期", error_code="AUTH_401_INVALID") from exc + return payload diff --git a/app/gateway/ownership.py b/app/gateway/ownership.py new file mode 100644 index 0000000..50f546b --- /dev/null +++ b/app/gateway/ownership.py @@ -0,0 +1,53 @@ +"""数据层归属校验(第二层)。""" + +from __future__ import annotations + +from app.gateway.jwt_service import permissions_for_roles +from app.model.schemas import AuthContext +from app.repository.advisor_rel_repository import AdvisorRelRepository +from app.repository.audit_repository import AuditRepository +from app.utils.exceptions import ForbiddenError + + +def audit_denied(ctx: AuthContext, error_code: str, customer_id: str | None) -> None: + AuditRepository().insert( + trace_id=ctx.trace_id, + event_type="auth_denied", + agent_type=ctx.agent_type, + actor_id=ctx.sub, + customer_id=customer_id, + decision=error_code, + input_summary={"roles": ctx.roles, "agent_type": ctx.agent_type}, + ) + + +def assert_customer_access(ctx: AuthContext, customer_id: str, *, action: str = "detail") -> None: + if ctx.token_type == "customer": + if ctx.sub != customer_id: + audit_denied(ctx, "AUTH_403_NOT_OWNER", customer_id) + raise ForbiddenError("无权访问该客户数据", error_code="AUTH_403_NOT_OWNER") + return + + if "advisor" in ctx.roles: + if not AdvisorRelRepository().is_assigned(ctx.sub, customer_id): + audit_denied(ctx, "AUTH_403_NOT_ASSIGNED", customer_id) + raise ForbiddenError("客户不在您的服务名下", error_code="AUTH_403_NOT_ASSIGNED") + return + + if "analyst" in ctx.roles: + if action == "detail" and "core:customer:read:detail" not in ctx.permissions: + audit_denied(ctx, "AUTH_403_SCOPE", customer_id) + raise ForbiddenError("分析员无客户明细权限", error_code="AUTH_403_SCOPE") + return + + if "risk_officer" in ctx.roles: + return + + audit_denied(ctx, "AUTH_403_ROLE", customer_id) + raise ForbiddenError("当前角色无权访问客户数据", error_code="AUTH_403_ROLE") + + +def resolve_effective_customer_id(ctx: AuthContext, requested: str | None) -> str | None: + if ctx.token_type == "customer": + return ctx.sub + return requested diff --git a/app/gateway/rbac.py b/app/gateway/rbac.py new file mode 100644 index 0000000..49f373f --- /dev/null +++ b/app/gateway/rbac.py @@ -0,0 +1,35 @@ +"""RBAC:Agent 入口与角色准入。""" + +from __future__ import annotations + +from app.model.schemas import AgentType, AuthContext +from app.utils.exceptions import ForbiddenError + +AGENT_ACCESS: dict[AgentType, dict[str, set[str]]] = { + "customer": {"token_types": {"customer"}, "roles": {"customer"}}, + "advisor": {"token_types": {"staff"}, "roles": {"advisor", "compliance", "ops"}}, + "analyst": {"token_types": {"staff"}, "roles": {"analyst", "compliance"}}, + "risk": {"token_types": {"staff", "service"}, "roles": {"risk_officer", "service_risk"}}, +} + + +def assert_agent_access(ctx: AuthContext) -> None: + rule = AGENT_ACCESS[ctx.agent_type] + if ctx.token_type not in rule["token_types"]: + raise ForbiddenError( + "Token 类型与 Agent 不匹配", + error_code="AUTH_403_AGENT_MISMATCH", + ) + if not rule["roles"].intersection(ctx.roles): + raise ForbiddenError( + "角色无权访问该 Agent", + error_code="AUTH_403_ROLE", + ) + perm = f"agent:{ctx.agent_type}:chat" + if ctx.agent_type == "risk" and ctx.token_type == "service": + return + if perm not in ctx.permissions and not ctx.has_perm(perm): + raise ForbiddenError( + "缺少 Agent 对话权限", + error_code="AUTH_403_PERMISSION", + ) diff --git a/app/main.py b/app/main.py index 12fbe63..ee4a147 100644 --- a/app/main.py +++ b/app/main.py @@ -1,15 +1,45 @@ """FastAPI 入口:挂载路由、中间件(JWT/RBAC)、生命周期。""" -from fastapi import FastAPI +from __future__ import annotations +from fastapi import FastAPI, Request +from fastapi.responses import JSONResponse + +from app.api.auth import router as auth_router +from app.api.chat import router as chat_router from app.config.settings import settings +from app.middleware.trace import TraceMiddleware +from app.repository.audit_repository import AuditRepository +from app.utils.exceptions import AppError +from app.utils.response import fail app = FastAPI(title="JinRong Agent Platform", version="0.1.0") +app.add_middleware(TraceMiddleware) +app.include_router(auth_router) +app.include_router(chat_router) @app.get("/health") -def health(): - return {"status": "ok", "env": settings.app_env} +def health(request: Request): + trace_id = getattr(request.state, "trace_id", "unknown") + return {"status": "ok", "env": settings.app_env, "trace_id": trace_id} -# TODO: 挂载 app.api 路由;接入 Auth SDK 中间件 +@app.exception_handler(AppError) +async def app_error_handler(request: Request, exc: AppError): + trace_id = getattr(request.state, "trace_id", "unknown") + if exc.audit_event: + agent_type = request.headers.get("X-Agent-Type", "platform") + try: + AuditRepository().insert( + trace_id=trace_id, + event_type=exc.audit_event, + agent_type=agent_type if agent_type in ("customer", "advisor", "analyst", "risk") else "platform", + actor_id="anonymous", + decision=exc.error_code, + input_summary={"message": exc.message}, + ) + except Exception: + pass + body = fail(exc.code, exc.message, trace_id, data={"error_code": exc.error_code}) + return JSONResponse(status_code=exc.http_status, content=body.model_dump()) diff --git a/app/middleware/__init__.py b/app/middleware/__init__.py new file mode 100644 index 0000000..7ff05be --- /dev/null +++ b/app/middleware/__init__.py @@ -0,0 +1 @@ +"""HTTP 中间件。""" diff --git a/app/middleware/trace.py b/app/middleware/trace.py new file mode 100644 index 0000000..11fcfac --- /dev/null +++ b/app/middleware/trace.py @@ -0,0 +1,18 @@ +"""全链路 trace_id 中间件。""" + +from __future__ import annotations + +import uuid + +from starlette.middleware.base import BaseHTTPMiddleware +from starlette.requests import Request +from starlette.responses import Response + + +class TraceMiddleware(BaseHTTPMiddleware): + async def dispatch(self, request: Request, call_next) -> Response: + trace_id = request.headers.get("X-Trace-Id") or str(uuid.uuid4()) + request.state.trace_id = trace_id + response = await call_next(request) + response.headers["X-Trace-Id"] = trace_id + return response diff --git a/app/model/schemas.py b/app/model/schemas.py index ffcf983..fc4719b 100644 --- a/app/model/schemas.py +++ b/app/model/schemas.py @@ -1 +1,67 @@ """Pydantic 请求/响应模型、AuthContext 等。""" + +from __future__ import annotations + +from typing import Any, Literal + +from pydantic import BaseModel, Field + +AgentType = Literal["customer", "advisor", "analyst", "risk"] +TokenType = Literal["customer", "staff", "service"] + + +class AuthContext(BaseModel): + """网关校验后注入业务层的身份上下文。""" + + sub: str + token_type: TokenType + roles: list[str] + permissions: list[str] = Field(default_factory=list) + tenant_id: str = "default" + trace_id: str + agent_type: AgentType + jti: str + customer_id: str | None = None + + def has_role(self, role: str) -> bool: + return role in self.roles + + def has_perm(self, perm: str) -> bool: + if perm in self.permissions: + return True + prefix = perm.split(":")[0] + return any(p.startswith(f"{prefix}:") and p.endswith(":*") for p in self.permissions) + + +class LoginRequest(BaseModel): + actor_id: str = Field(..., description="customer_id 或 staff_id") + token_type: TokenType = "staff" + roles: list[str] | None = Field(default=None, description="开发期可覆盖角色,默认按 actor_id 推断") + + +class LoginResponseData(BaseModel): + access_token: str + token_type: str = "Bearer" + expires_in: int + sub: str + roles: list[str] + + +class ChatRequest(BaseModel): + message: str = Field(..., min_length=1, max_length=8000) + session_id: str | None = None + customer_id: str | None = Field(default=None, description="员工 Agent 查询目标客户时使用") + + +class ChatResponseData(BaseModel): + session_id: str + reply: str + agent_type: AgentType + has_disclaimer: bool = False + + +class ApiResponse(BaseModel): + code: int = 0 + message: str = "ok" + data: Any = None + trace_id: str diff --git a/app/repository/advisor_rel_repository.py b/app/repository/advisor_rel_repository.py new file mode 100644 index 0000000..2cb877a --- /dev/null +++ b/app/repository/advisor_rel_repository.py @@ -0,0 +1,31 @@ +"""客户-代理人归属(jinrong_agent.customer_advisor_rel)。""" + +from __future__ import annotations + +from sqlalchemy import text +from sqlalchemy.engine import Engine + +from app.config.database import get_agent_engine +from app.repository.core_ro import CoreReadOnlyRepository + + +class AdvisorRelRepository: + def __init__(self, engine: Engine | None = None, core_ro: CoreReadOnlyRepository | None = None) -> None: + self._engine = engine or get_agent_engine() + self._core = core_ro or CoreReadOnlyRepository() + + def is_assigned(self, advisor_id: str, customer_id: str) -> bool: + sql = text( + """ + SELECT 1 FROM customer_advisor_rel + WHERE advisor_id = :aid AND customer_id = :cid AND rel_status = 'active' + LIMIT 1 + """ + ) + try: + with self._engine.connect() as conn: + if conn.execute(sql, {"aid": advisor_id, "cid": customer_id}).first(): + return True + except Exception: + pass + return self._core.is_advisor_assigned(advisor_id, customer_id) diff --git a/app/repository/agent_repository.py b/app/repository/agent_repository.py new file mode 100644 index 0000000..2e9d154 --- /dev/null +++ b/app/repository/agent_repository.py @@ -0,0 +1,90 @@ +"""Agent 会话与消息持久化。""" + +from __future__ import annotations + +import uuid + +from sqlalchemy import text +from sqlalchemy.engine import Engine + +from app.config.database import get_agent_engine +from app.model.schemas import AgentType, AuthContext + + +class AgentSessionRepository: + def __init__(self, engine: Engine | None = None) -> None: + self._engine = engine or get_agent_engine() + + def ensure_session( + self, + ctx: AuthContext, + session_id: str | None, + customer_id: str | None, + ) -> str: + sid = session_id or str(uuid.uuid4()) + sql = text( + """ + INSERT INTO agent_session + (session_id, trace_id, agent_type, actor_id, actor_role, + customer_id, advisor_id, status) + VALUES + (:session_id, :trace_id, :agent_type, :actor_id, :actor_role, + :customer_id, :advisor_id, 'active') + ON DUPLICATE KEY UPDATE + trace_id = VALUES(trace_id), + updated_at = CURRENT_TIMESTAMP(3) + """ + ) + actor_role = ctx.roles[0] if ctx.roles else "unknown" + advisor_id = ctx.sub if "advisor" in ctx.roles else None + with self._engine.begin() as conn: + conn.execute( + sql, + { + "session_id": sid, + "trace_id": ctx.trace_id, + "agent_type": ctx.agent_type, + "actor_id": ctx.sub, + "actor_role": actor_role, + "customer_id": customer_id, + "advisor_id": advisor_id, + }, + ) + return sid + + def next_seq(self, session_id: str) -> int: + sql = text("SELECT COALESCE(MAX(seq_no), 0) + 1 AS n FROM agent_message WHERE session_id = :sid") + with self._engine.connect() as conn: + row = conn.execute(sql, {"sid": session_id}).mappings().first() + return int(row["n"]) if row else 1 + + def insert_message( + self, + *, + session_id: str, + trace_id: str, + seq_no: int, + role: str, + content: str, + has_disclaimer: bool = False, + ) -> None: + sql = text( + """ + INSERT INTO agent_message + (session_id, trace_id, seq_no, role, content, has_disclaimer) + VALUES + (:session_id, :trace_id, :seq_no, :role, :content, :has_disclaimer) + """ + ) + with self._engine.begin() as conn: + conn.execute( + sql, + { + "session_id": session_id, + "trace_id": trace_id, + "seq_no": seq_no, + "role": role, + "content": content, + "has_disclaimer": 1 if has_disclaimer else 0, + }, + ) diff --git a/app/repository/audit_repository.py b/app/repository/audit_repository.py new file mode 100644 index 0000000..8b3b36a --- /dev/null +++ b/app/repository/audit_repository.py @@ -0,0 +1,52 @@ +"""审计总账写入(只 INSERT)。""" + +from __future__ import annotations + +import json +from typing import Any + +from sqlalchemy import text +from sqlalchemy.engine import Engine + +from app.config.database import get_agent_engine +from app.model.schemas import AgentType + + +class AuditRepository: + def __init__(self, engine: Engine | None = None) -> None: + self._engine = engine or get_agent_engine() + + def insert( + self, + *, + trace_id: str, + event_type: str, + agent_type: AgentType | str, + actor_id: str, + customer_id: str | None = None, + rule_id: str | None = None, + input_summary: dict[str, Any] | None = None, + decision: str | None = None, + ) -> None: + sql = text( + """ + INSERT INTO audit_log + (trace_id, event_type, agent_type, actor_id, customer_id, + rule_id, input_summary, decision) + VALUES + (:trace_id, :event_type, :agent_type, :actor_id, :customer_id, + :rule_id, :input_summary, :decision) + """ + ) + payload = { + "trace_id": trace_id, + "event_type": event_type, + "agent_type": agent_type if agent_type != "platform" else "platform", + "actor_id": actor_id, + "customer_id": customer_id, + "rule_id": rule_id, + "input_summary": json.dumps(input_summary, ensure_ascii=False) if input_summary else None, + "decision": decision, + } + with self._engine.begin() as conn: + conn.execute(sql, payload) diff --git a/app/service/agent_service.py b/app/service/agent_service.py index 0e607f8..63fc096 100644 --- a/app/service/agent_service.py +++ b/app/service/agent_service.py @@ -1 +1,64 @@ -"""Agent 编排:LangGraph StateGraph + DeepSeek;Tool 节点;四 Agent 能力边界。""" +"""Agent 编排:LangGraph StateGraph 最小骨架。""" + +from __future__ import annotations + +from typing import TypedDict + +from langgraph.graph import END, StateGraph + +from app.model.schemas import AgentType, AuthContext + +DISCLAIMER = ( + "本内容仅为投资分析参考,不构成任何直接投资建议,不构成对任何产品的收益承诺," + "据此操作风险自负,请谨慎对待。" +) + + +class AgentState(TypedDict): + message: str + reply: str + agent_type: AgentType + actor_id: str + customer_id: str | None + + +def _build_graph(): + graph = StateGraph(AgentState) + + def respond(state: AgentState) -> AgentState: + agent = state["agent_type"] + prefix = { + "customer": "客户财富助手", + "advisor": "代理人助手", + "analyst": "数据分析助手", + "risk": "风控监测助手", + }.get(agent, "助手") + target = f"(客户 {state['customer_id']})" if state.get("customer_id") else "" + reply = f"【{prefix}】{target}已收到:{state['message']}" + if agent in ("advisor", "analyst") and any( + k in state["message"] for k in ("报告", "建议", "推荐") + ): + reply = f"{reply}\n\n{DISCLAIMER}" + return {"reply": reply} + + graph.add_node("respond", respond) + graph.set_entry_point("respond") + graph.add_edge("respond", END) + return graph.compile() + + +_GRAPH = _build_graph() + + +def run_chat(ctx: AuthContext, message: str, customer_id: str | None) -> tuple[str, bool]: + state: AgentState = { + "message": message, + "reply": "", + "agent_type": ctx.agent_type, + "actor_id": ctx.sub, + "customer_id": customer_id, + } + result = _GRAPH.invoke(state) + reply = result["reply"] + has_disclaimer = DISCLAIMER in reply + return reply, has_disclaimer diff --git a/app/service/memory_service.py b/app/service/memory_service.py index 44ee227..ecb63bd 100644 --- a/app/service/memory_service.py +++ b/app/service/memory_service.py @@ -1 +1,13 @@ -"""记忆服务:Redis 会话窗口 + MySQL 落盘;L1/L2/L3 画像读写。""" +"""会话短期记忆薄封装(Wave 0:接口占位,后续接 Redis)。""" + +from __future__ import annotations + +from app.model.schemas import AuthContext + + +class MemoryService: + def session_key(self, session_id: str) -> str: + return f"session:{session_id}:messages" + + def touch_session(self, ctx: AuthContext, session_id: str) -> None: + _ = (ctx, session_id) diff --git a/app/utils/exceptions.py b/app/utils/exceptions.py index a7f5d40..80ccdd9 100644 --- a/app/utils/exceptions.py +++ b/app/utils/exceptions.py @@ -1 +1,37 @@ -"""业务异常与鉴权错误码(对齐 JWT-RBAC 手册)。""" +"""统一业务异常。""" + +from __future__ import annotations + + +class AppError(Exception): + def __init__( + self, + code: int, + message: str, + *, + http_status: int = 400, + error_code: str | None = None, + audit_event: str | None = None, + ) -> None: + super().__init__(message) + self.code = code + self.message = message + self.http_status = http_status + self.error_code = error_code + self.audit_event = audit_event + + +class UnauthorizedError(AppError): + def __init__(self, message: str = "未授权", *, error_code: str = "AUTH_401") -> None: + super().__init__(401, message, http_status=401, error_code=error_code, audit_event="auth_denied") + + +class ForbiddenError(AppError): + def __init__( + self, + message: str = "无权访问", + *, + error_code: str = "AUTH_403", + audit_event: str = "auth_denied", + ) -> None: + super().__init__(403, message, http_status=403, error_code=error_code, audit_event=audit_event) diff --git a/app/utils/input_guard.py b/app/utils/input_guard.py new file mode 100644 index 0000000..8dd889d --- /dev/null +++ b/app/utils/input_guard.py @@ -0,0 +1,25 @@ +"""输入防护(F-03 最小实现)。""" + +from __future__ import annotations + +import re + +from app.utils.exceptions import AppError + +_INJECTION_PATTERNS = ( + re.compile(r"(?i)ignore\s+previous\s+instructions"), + re.compile(r"(?i)system\s*:\s*"), + re.compile(r"(?is)(drop|delete|update|insert|alter|truncate)\s+"), +) + + +def validate_user_message(message: str, *, max_len: int = 8000) -> str: + text = message.strip() + if not text: + raise AppError(400, "消息不能为空", error_code="INPUT_EMPTY") + if len(text) > max_len: + raise AppError(400, "消息过长", error_code="INPUT_OVERSIZE", audit_event="input_guard") + for pattern in _INJECTION_PATTERNS: + if pattern.search(text): + raise AppError(400, "输入包含不允许的内容", error_code="INPUT_BLOCKED", audit_event="input_guard") + return text diff --git a/app/utils/logger.py b/app/utils/logger.py index 3bc6aa2..b3b4488 100644 --- a/app/utils/logger.py +++ b/app/utils/logger.py @@ -1 +1,6 @@ -"""日志模块。""" +"""结构化日志(不打印完整 JWT)。""" + +import logging + +logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s %(message)s") +logger = logging.getLogger("jinrong") diff --git a/app/utils/response.py b/app/utils/response.py index a868c47..8efdc78 100644 --- a/app/utils/response.py +++ b/app/utils/response.py @@ -1 +1,15 @@ """统一 API 响应格式。""" + +from __future__ import annotations + +from typing import Any + +from app.model.schemas import ApiResponse + + +def ok(data: Any, trace_id: str, message: str = "ok") -> ApiResponse: + return ApiResponse(code=0, message=message, data=data, trace_id=trace_id) + + +def fail(code: int, message: str, trace_id: str, *, data: Any = None) -> ApiResponse: + return ApiResponse(code=code, message=message, data=data, trace_id=trace_id) diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..dfe23e2 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,42 @@ +"""pytest 公共 fixture。""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest +from fastapi.testclient import TestClient + +from app.main import app +from app.repository.advisor_rel_repository import AdvisorRelRepository +from app.repository.agent_repository import AgentSessionRepository +from app.repository.audit_repository import AuditRepository + + +@pytest.fixture +def client(): + return TestClient(app) + + +@pytest.fixture(autouse=True) +def mock_db(monkeypatch): + audit = MagicMock() + session = MagicMock() + advisor = MagicMock() + + audit.insert = MagicMock() + session.ensure_session = MagicMock(return_value="sess-test-001") + session.next_seq = MagicMock(return_value=1) + session.insert_message = MagicMock() + advisor.is_assigned = MagicMock(return_value=True) + + monkeypatch.setattr(AuditRepository, "__init__", lambda self, engine=None: None) + monkeypatch.setattr(AuditRepository, "insert", audit.insert) + monkeypatch.setattr(AgentSessionRepository, "__init__", lambda self, engine=None: None) + monkeypatch.setattr(AgentSessionRepository, "ensure_session", session.ensure_session) + monkeypatch.setattr(AgentSessionRepository, "next_seq", session.next_seq) + monkeypatch.setattr(AgentSessionRepository, "insert_message", session.insert_message) + monkeypatch.setattr(AdvisorRelRepository, "__init__", lambda self, engine=None, core_ro=None: None) + monkeypatch.setattr(AdvisorRelRepository, "is_assigned", advisor.is_assigned) + + yield {"audit": audit, "session": session, "advisor": advisor} diff --git a/tests/test_wave0_auth.py b/tests/test_wave0_auth.py new file mode 100644 index 0000000..37afdf5 --- /dev/null +++ b/tests/test_wave0_auth.py @@ -0,0 +1,78 @@ +"""Wave 0:JWT / RBAC 测试。""" + +from app.gateway.jwt_service import decode_token, issue_token + + +def test_issue_and_decode_staff_token(): + token, expires = issue_token("STAFF-20001", "staff") + assert expires > 0 + payload = decode_token(token) + assert payload["sub"] == "STAFF-20001" + assert "analyst" in payload["roles"] + assert "agent:analyst:chat" in payload["permissions"] + + +def test_login_endpoint(client): + resp = client.post("/api/auth/login", json={"actor_id": "STAFF-20001", "token_type": "staff"}) + assert resp.status_code == 200 + body = resp.json() + assert body["code"] == 0 + assert "access_token" in body["data"] + assert "trace_id" in body + + +def test_chat_requires_auth(client): + resp = client.post( + "/api/chat", + json={"message": "hello"}, + headers={"X-Agent-Type": "analyst"}, + ) + assert resp.status_code == 401 + + +def test_chat_analyst_ok(client): + login = client.post("/api/auth/login", json={"actor_id": "STAFF-20001", "token_type": "staff"}) + token = login.json()["data"]["access_token"] + resp = client.post( + "/api/chat", + json={"message": "上季度收益率"}, + headers={ + "Authorization": f"Bearer {token}", + "X-Agent-Type": "analyst", + }, + ) + assert resp.status_code == 200 + body = resp.json() + assert body["code"] == 0 + assert "reply" in body["data"] + assert body["data"]["agent_type"] == "analyst" + + +def test_agent_type_mismatch_forbidden(client): + login = client.post("/api/auth/login", json={"actor_id": "STAFF-20001", "token_type": "staff"}) + token = login.json()["data"]["access_token"] + resp = client.post( + "/api/chat", + json={"message": "hello"}, + headers={ + "Authorization": f"Bearer {token}", + "X-Agent-Type": "customer", + }, + ) + assert resp.status_code == 403 + + +def test_advisor_not_assigned_forbidden(client, mock_db): + mock_db["advisor"].is_assigned.return_value = False + login = client.post("/api/auth/login", json={"actor_id": "STAFF-10086", "token_type": "staff"}) + token = login.json()["data"]["access_token"] + resp = client.post( + "/api/chat", + json={"message": "查持仓", "customer_id": "CUST-1010"}, + headers={ + "Authorization": f"Bearer {token}", + "X-Agent-Type": "advisor", + }, + ) + assert resp.status_code == 403 + assert mock_db["audit"].insert.called diff --git a/tests/test_wave0_input_guard.py b/tests/test_wave0_input_guard.py new file mode 100644 index 0000000..c3628d9 --- /dev/null +++ b/tests/test_wave0_input_guard.py @@ -0,0 +1,16 @@ +"""Wave 0:输入防护测试。""" + + +def test_block_sql_injection_in_chat(client): + login = client.post("/api/auth/login", json={"actor_id": "STAFF-20001", "token_type": "staff"}) + token = login.json()["data"]["access_token"] + resp = client.post( + "/api/chat", + json={"message": "please DROP TABLE users"}, + headers={ + "Authorization": f"Bearer {token}", + "X-Agent-Type": "analyst", + }, + ) + assert resp.status_code == 400 + assert resp.json()["code"] == 400 diff --git a/tests/test_wave0_trace.py b/tests/test_wave0_trace.py new file mode 100644 index 0000000..0120aab --- /dev/null +++ b/tests/test_wave0_trace.py @@ -0,0 +1,13 @@ +"""Wave 0:trace_id 测试。""" + + +def test_health_returns_trace_id(client): + resp = client.get("/health") + assert resp.status_code == 200 + assert resp.headers.get("X-Trace-Id") + assert resp.json()["trace_id"] + + +def test_trace_id_echo(client): + resp = client.get("/health", headers={"X-Trace-Id": "trace-fixed-001"}) + assert resp.headers["X-Trace-Id"] == "trace-fixed-001"