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:
2026-09-07 17:20:42 +08:00
parent 5f0f86a007
commit 3995cb44d8
26 changed files with 1013 additions and 10 deletions
+25
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+4
View File
@@ -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()
+1
View File
@@ -0,0 +1 @@
"""Gateway 模块。"""
+55
View File
@@ -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
+143
View File
@@ -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
+53
View File
@@ -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
+35
View File
@@ -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
View File
@@ -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())
+1
View File
@@ -0,0 +1 @@
"""HTTP 中间件。"""
+18
View File
@@ -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
+66
View File
@@ -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
+31
View File
@@ -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)
+90
View File
@@ -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,
},
)
+52
View File
@@ -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)
+64 -1
View File
@@ -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
+13 -1
View File
@@ -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
View File
@@ -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)
+25
View File
@@ -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
View File
@@ -1 +1,6 @@
"""日志模块。"""
"""结构化日志(不打印完整 JWT)。"""
import logging
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s %(message)s")
logger = logging.getLogger("jinrong")
+14
View File
@@ -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)
+42
View File
@@ -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}
+78
View File
@@ -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
+16
View File
@@ -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
+13
View File
@@ -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"