Files
X-Financial/server/tests/test_orchestrator_auth_endpoints.py

408 lines
14 KiB
Python
Raw Permalink Normal View History

from __future__ import annotations
from collections.abc import Generator
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from sqlalchemy import create_engine, select
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.api.deps import (
CurrentUserContext,
get_current_user,
get_db,
get_optional_current_user,
)
from app.db.base import Base
from app.main import create_app
from app.models.agent_conversation import AgentConversation
from app.schemas.orchestrator import (
OrchestratorRequest,
OrchestratorResponse,
OrchestratorTraceSummary,
)
from app.services.orchestrator import OrchestratorService
@pytest.fixture
def http_context() -> Generator[
tuple[TestClient, FastAPI, sessionmaker[Session]],
None,
None,
]:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = create_app()
def override_db() -> Generator[Session, None, None]:
with session_factory() as db:
yield db
app.dependency_overrides[get_db] = override_db
client = TestClient(app)
try:
yield client, app, session_factory
finally:
client.close()
app.dependency_overrides.clear()
engine.dispose()
def authenticated_user() -> CurrentUserContext:
return CurrentUserContext(
username="signed-in@example.com",
name="登录用户",
role_codes=["user"],
is_admin=False,
tenant_id="tenant-signed-in",
department_name="财务共享部",
cost_center="CC-SIGNED-IN",
position="费用专员",
grade="P5",
employee_no="E-SIGNED-IN",
manager_name="直属经理",
employee_id="employee-signed-in",
auth_session_id="session-signed-in",
)
def install_authenticated_user(app: FastAPI) -> CurrentUserContext:
current_user = authenticated_user()
app.dependency_overrides[get_current_user] = lambda: current_user
app.dependency_overrides[get_optional_current_user] = lambda: current_user
return current_user
def install_admin_user(app: FastAPI) -> CurrentUserContext:
current_user = authenticated_user()
current_user.role_codes = ["admin"]
current_user.is_admin = True
app.dependency_overrides[get_current_user] = lambda: current_user
app.dependency_overrides[get_optional_current_user] = lambda: current_user
return current_user
@pytest.mark.parametrize("source", ["user_message", "system_event", "schedule"])
def test_orchestrator_run_requires_authentication_for_every_external_source(
http_context: tuple[TestClient, FastAPI, sessionmaker[Session]],
source: str,
) -> None:
client, _, _ = http_context
response = client.post(
"/api/v1/orchestrator/run",
json={
"source": source,
"user_id": "forged-user@example.com",
"message": "帮我申请差旅费用",
"context_json": {},
},
)
assert response.status_code == 401
assert response.headers["www-authenticate"] == "Bearer"
assert response.json()["detail"] == "请先登录后再使用智能助手。"
@pytest.mark.parametrize("source", ["system_event", "schedule"])
def test_non_admin_cannot_trigger_privileged_orchestrator_sources(
http_context: tuple[TestClient, FastAPI, sessionmaker[Session]],
source: str,
) -> None:
client, app, _ = http_context
install_authenticated_user(app)
response = client.post(
"/api/v1/orchestrator/run",
json={
"source": source,
"user_id": "forged-admin@example.com",
"message": "运行全局风险扫描",
"context_json": {
"role_codes": ["admin"],
"is_admin": True,
},
},
)
assert response.status_code == 403
assert response.json()["detail"] == "只有平台管理员可以从外部触发调度或系统事件。"
@pytest.mark.parametrize("source", ["system_event", "schedule"])
def test_admin_can_trigger_privileged_orchestrator_sources(
http_context: tuple[TestClient, FastAPI, sessionmaker[Session]],
monkeypatch: pytest.MonkeyPatch,
source: str,
) -> None:
client, app, _ = http_context
expected_user = install_admin_user(app)
captured: dict[str, object] = {}
def fake_run(
self: OrchestratorService,
payload: OrchestratorRequest,
*,
current_user: CurrentUserContext | None = None,
) -> OrchestratorResponse:
del self
captured["source"] = payload.source
captured["current_user"] = current_user
return OrchestratorResponse(
run_id="run-admin-source",
selected_agent="hermes",
route_reason="test_admin_source",
permission_level="read",
status="succeeded",
result={},
requires_confirmation=False,
trace_summary=OrchestratorTraceSummary(
scenario="system",
intent="risk_check",
),
)
monkeypatch.setattr(OrchestratorService, "run", fake_run)
response = client.post(
"/api/v1/orchestrator/run",
json={
"source": source,
"message": "运行全局风险扫描",
"context_json": {},
},
)
assert response.status_code == 200, response.text
assert captured["source"] == source
assert captured["current_user"] == expected_user
@pytest.mark.parametrize("source", ["user_message", "system_event", "schedule"])
def test_authenticated_payload_overrides_all_scheduler_identity_aliases(
source: str,
) -> None:
current_user = authenticated_user()
payload = OrchestratorRequest(
source=source,
user_id="forged-user@example.com",
message="测试身份覆盖",
context_json={
"username": "forged-user@example.com",
"user_id": "forged-user@example.com",
"tenant_id": "tenant-forged",
"role_codes": ["admin"],
"is_admin": True,
"requested_by_username": "forged-requester@example.com",
"requested_by_name": "伪造请求人",
"actor": "forged-actor@example.com",
"actor_id": "forged-actor-id",
"auth_session_id": "forged-session",
},
)
trusted = OrchestratorService._build_authenticated_user_payload(payload, current_user)
assert trusted.user_id == current_user.username
assert trusted.context_json["tenant_id"] == current_user.tenant_id
assert trusted.context_json["role_codes"] == current_user.role_codes
assert trusted.context_json["is_admin"] is current_user.is_admin
assert trusted.context_json["requested_by_username"] == current_user.username
assert trusted.context_json["requested_by_name"] == current_user.name
assert trusted.context_json["actor"] == current_user.username
assert trusted.context_json["actor_id"] == current_user.username
assert "auth_session_id" not in trusted.context_json
def test_authenticated_run_passes_server_identity_to_orchestrator(
http_context: tuple[TestClient, FastAPI, sessionmaker[Session]],
monkeypatch: pytest.MonkeyPatch,
) -> None:
client, app, _ = http_context
expected_user = install_authenticated_user(app)
captured: dict[str, object] = {}
def fake_run(
self: OrchestratorService,
payload: OrchestratorRequest,
*,
current_user: CurrentUserContext | None = None,
) -> OrchestratorResponse:
del self
captured["payload"] = payload
captured["current_user"] = current_user
return OrchestratorResponse(
run_id="run-auth-binding",
conversation_id="conversation-auth-binding",
selected_agent="user_agent",
route_reason="test",
permission_level="user",
status="succeeded",
result={},
requires_confirmation=False,
trace_summary=OrchestratorTraceSummary(
scenario="expense",
intent="operate",
),
)
monkeypatch.setattr(OrchestratorService, "run", fake_run)
response = client.post(
"/api/v1/orchestrator/run",
json={
"source": "user_message",
"user_id": "forged-user@example.com",
"message": "保存申请草稿",
"context_json": {
"username": "forged-user@example.com",
"tenant_id": "tenant-forged",
"role_codes": ["admin"],
"is_admin": True,
"employee_no": "E-FORGED",
"auth_session_id": "session-forged",
},
},
)
assert response.status_code == 200, response.text
assert captured["current_user"] == expected_user
assert captured["payload"].user_id == "forged-user@example.com"
assert captured["payload"].context_json["tenant_id"] == "tenant-forged"
def test_latest_conversation_uses_authenticated_username(
http_context: tuple[TestClient, FastAPI, sessionmaker[Session]],
) -> None:
client, app, session_factory = http_context
current_user = install_authenticated_user(app)
with session_factory() as db:
db.add_all(
[
AgentConversation(
conversation_id="conversation-signed-in",
user_id=current_user.username,
source="user_message",
title="登录用户会话",
state_json={
"tenant_id": current_user.tenant_id,
"session_type": "expense",
},
),
AgentConversation(
conversation_id="conversation-same-user-other-tenant",
user_id=current_user.username,
source="user_message",
title="其他租户同名用户会话",
state_json={
"tenant_id": "tenant-other",
"session_type": "expense",
"application_preview_decision": {
"decision_id": "decision-other-tenant",
"application_preview": {
"fields": {"reason": "其他租户敏感项目"}
},
},
},
),
AgentConversation(
conversation_id="conversation-forged-user",
user_id="forged-user@example.com",
source="user_message",
title="伪造用户会话",
state_json={
"tenant_id": "tenant-forged",
"session_type": "expense",
},
),
]
)
db.commit()
response = client.get(
"/api/v1/orchestrator/conversations/latest",
params={"user_id": "forged-user@example.com", "session_type": "expense"},
)
assert response.status_code == 200, response.text
payload = response.json()
assert payload["found"] is True
assert payload["conversation"]["conversation_id"] == "conversation-signed-in"
assert payload["conversation"]["user_id"] == current_user.username
def test_delete_conversations_uses_authenticated_username(
http_context: tuple[TestClient, FastAPI, sessionmaker[Session]],
) -> None:
client, app, session_factory = http_context
current_user = install_authenticated_user(app)
with session_factory() as db:
db.add_all(
[
AgentConversation(
conversation_id="conversation-owner-single",
user_id=current_user.username,
source="user_message",
state_json={
"tenant_id": current_user.tenant_id,
"session_type": "expense",
},
),
AgentConversation(
conversation_id="conversation-owner-bulk",
user_id=current_user.username,
source="user_message",
state_json={
"tenant_id": current_user.tenant_id,
"session_type": "expense",
},
),
AgentConversation(
conversation_id="conversation-same-user-other-tenant",
user_id=current_user.username,
source="user_message",
state_json={
"tenant_id": "tenant-other",
"session_type": "expense",
},
),
AgentConversation(
conversation_id="conversation-attacker",
user_id="forged-user@example.com",
source="user_message",
state_json={
"tenant_id": "tenant-forged",
"session_type": "expense",
},
),
]
)
db.commit()
single_response = client.delete(
"/api/v1/orchestrator/conversations/conversation-owner-single",
params={"user_id": "forged-user@example.com"},
)
bulk_response = client.delete(
"/api/v1/orchestrator/conversations",
params={"user_id": "forged-user@example.com", "session_type": "expense"},
)
assert single_response.status_code == 200, single_response.text
assert single_response.json()["deleted_count"] == 1
assert bulk_response.status_code == 200, bulk_response.text
assert bulk_response.json()["deleted_count"] == 1
with session_factory() as db:
remaining = list(db.scalars(select(AgentConversation)).all())
assert {item.conversation_id for item in remaining} == {
"conversation-attacker",
"conversation-same-user-other-tenant",
}