Implement authentication and chat functionality with JWT support
- Added `auth.py` for mock login and JWT issuance. - Introduced `chat.py` for handling chat requests with role-based access control. - Enhanced `main.py` to include new routers and middleware for tracing. - Implemented input validation in `input_guard.py` to prevent SQL injection. - Created repositories for managing agent sessions and audit logs. - Added exception handling for authorization errors. - Updated settings to include JWT configuration. - Introduced tests for authentication and input validation.
This commit is contained in:
@@ -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")
|
||||
+70
-1
@@ -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)
|
||||
|
||||
+27
-1
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Gateway 模块。"""
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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",
|
||||
)
|
||||
+34
-4
@@ -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())
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""HTTP 中间件。"""
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
},
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
+37
-1
@@ -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)
|
||||
|
||||
@@ -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
|
||||
+6
-1
@@ -1 +1,6 @@
|
||||
"""日志模块。"""
|
||||
"""结构化日志(不打印完整 JWT)。"""
|
||||
|
||||
import logging
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s %(message)s")
|
||||
logger = logging.getLogger("jinrong")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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}
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user