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

415 lines
14 KiB
Python
Raw Normal View History

from __future__ import annotations
from collections.abc import Generator
from datetime import UTC, datetime
from pathlib import Path
from types import MethodType
import pytest
from auth_helpers import install_legacy_header_auth_override
from fastapi import FastAPI
from fastapi.testclient import TestClient
from sqlalchemy import create_engine, event, select
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.api.deps import CurrentUserContext, get_db
from app.api.v1.endpoints.agent_asset_releases import _actor as release_actor
from app.api.v1.endpoints.agent_assets import router as agent_assets_router
from app.db.base import Base
from app.models.agent_asset import AgentAsset, AgentAssetTestRun, AgentAssetVersion
from app.models.financial_record import ExpenseClaim
from app.models.tenant import Tenant
from app.schemas.agent_asset import AgentAssetRiskRuleScenarioTestRequest
from app.services.agent_asset_access import stable_user_principal
from app.services.agent_asset_onlyoffice_security import (
AgentAssetOnlyOfficeSessionService,
)
from app.services.agent_assets import AgentAssetService
from app.services.settings import OnlyOfficeRuntimeConfig
def _factory() -> sessionmaker[Session]:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(engine)
return sessionmaker(bind=engine, autoflush=False, autocommit=False)
def _seed(db: Session) -> dict[str, AgentAsset]:
db.add_all(
[
Tenant(tenant_id="platform", tenant_code="platform", name="平台管理域"),
Tenant(tenant_id="tenant-a", tenant_code="tenant-a", name="租户 A"),
Tenant(tenant_id="tenant-b", tenant_code="tenant-b", name="租户 B"),
]
)
assets = {
"platform": AgentAsset(
tenant_id="platform",
scope="platform",
asset_type="rule",
code="rule.shared",
name="平台规则",
domain="expense",
owner="platform",
status="draft",
current_version="v1.0.0",
working_version="v1.0.0",
),
"tenant-a": AgentAsset(
tenant_id="tenant-a",
scope="tenant",
asset_type="rule",
code="rule.tenant",
name="租户 A 规则",
domain="expense",
owner="tenant-a",
status="draft",
current_version="v1.0.0",
working_version="v1.0.0",
),
"tenant-b": AgentAsset(
tenant_id="tenant-b",
scope="tenant",
asset_type="rule",
code="rule.tenant",
name="租户 B 规则",
domain="expense",
owner="tenant-b",
status="draft",
current_version="v1.0.0",
working_version="v1.0.0",
),
}
db.add_all(assets.values())
db.flush()
for asset in assets.values():
db.add(
AgentAssetVersion(
tenant_id=asset.tenant_id,
scope=asset.scope,
asset_id=asset.id,
version="v1.0.0",
content="# v1",
content_type="markdown",
created_by="seed",
)
)
db.commit()
return assets
def _client(factory: sessionmaker[Session], *, auth_override: bool = True) -> TestClient:
app = FastAPI()
app.include_router(agent_assets_router)
if auth_override:
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
with factory() as db:
yield db
app.dependency_overrides[get_db] = override_db
return TestClient(app)
def _headers(
tenant_id: str,
*,
employee_id: str = "employee-a",
is_admin: bool = False,
) -> dict[str, str]:
return {
"X-Auth-Username": f"user-{tenant_id}",
"X-Auth-Name": "Mutable Display Name",
"X-Auth-Employee-Id": employee_id,
"X-Auth-Role-Codes": "finance,manager",
"X-Auth-Is-Admin": "true" if is_admin else "false",
"X-Auth-Tenant-Id": tenant_id,
}
@pytest.fixture(autouse=True)
def _skip_foundation_bootstrap(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(
"app.services.agent_foundation.AgentFoundationService.ensure_foundation_ready",
lambda _self: None,
)
def test_asset_reads_are_authenticated_and_tenant_scoped() -> None:
factory = _factory()
with factory() as db:
assets = _seed(db)
ids = {key: value.id for key, value in assets.items()}
client = _client(factory)
listing = client.get("/agent-assets", headers=_headers("tenant-a"))
assert listing.status_code == 200
visible_ids = {item["id"] for item in listing.json()}
assert ids["platform"] in visible_ids
assert ids["tenant-a"] in visible_ids
assert ids["tenant-b"] not in visible_ids
assert (
client.get(
f"/agent-assets/{ids['tenant-b']}", headers=_headers("tenant-a")
).status_code
== 404
)
assert (
client.get(
f"/agent-assets/{ids['tenant-b']}/versions",
headers=_headers("tenant-a"),
).status_code
== 404
)
anonymous = _client(factory, auth_override=False).get("/agent-assets")
assert anonymous.status_code == 401
def test_version_write_uses_stable_principal_and_blocks_cross_tenant_or_platform() -> None:
factory = _factory()
with factory() as db:
assets = _seed(db)
ids = {key: value.id for key, value in assets.items()}
client = _client(factory)
body = {
"version": "v1.0.1",
"content": "# tenant update",
"content_type": "markdown",
"created_by": "spoofed-display-name",
}
cross_tenant = client.post(
f"/agent-assets/{ids['tenant-b']}/versions",
json=body,
headers=_headers("tenant-a"),
)
assert cross_tenant.status_code == 404
platform_by_editor = client.post(
f"/agent-assets/{ids['platform']}/versions",
json=body,
headers=_headers("tenant-a"),
)
assert platform_by_editor.status_code == 400
own = client.post(
f"/agent-assets/{ids['tenant-a']}/versions",
json=body,
headers=_headers("tenant-a", employee_id="employee-stable"),
)
assert own.status_code == 201
assert own.json()["created_by"] == "employee:employee-stable"
assert own.json()["tenant_id"] == "tenant-a"
assert own.json()["scope"] == "tenant"
platform_admin = client.post(
f"/agent-assets/{ids['platform']}/versions",
json=body,
headers=_headers("tenant-a", is_admin=True),
)
assert platform_admin.status_code == 201
assert platform_admin.json()["scope"] == "platform"
def test_scenario_samples_filter_tenant_in_sql_and_persist_target_tenant() -> None:
factory = _factory()
with factory() as db:
assets = _seed(db)
platform_asset = assets["platform"]
db.add_all(
[
_claim("tenant-a", "A-001"),
_claim("tenant-b", "B-001"),
]
)
db.commit()
statements: list[str] = []
@event.listens_for(db.get_bind(), "before_cursor_execute")
def _capture(_conn, _cursor, statement, _parameters, _context, _executemany):
if "FROM expense_claims" in statement:
statements.append(statement)
current_user = CurrentUserContext(
username="reviewer-a",
name="显示名一",
role_codes=["manager"],
is_admin=True,
tenant_id="tenant-a",
employee_id="employee-reviewer-a",
)
service = AgentAssetService(db, current_user=current_user)
def _load(_self, _asset_id: str, _version: str | None):
return platform_asset, "v1.0.0", {}
def _run(_self, _manifest: dict, claim: ExpenseClaim):
return {"claim_id": claim.id, "hit": False, "severity": "none"}
service._load_risk_rule_for_test = MethodType(_load, service)
service._run_claim_scenario = MethodType(_run, service)
result = service.run_risk_rule_scenario_test(
platform_asset.id,
AgentAssetRiskRuleScenarioTestRequest(
target_tenant_id="tenant-a",
intent="最近 30 天",
),
actor="employee:employee-reviewer-a",
)
assert result.result_json["total_count"] == 1
assert statements
assert "expense_claims.tenant_id = ?" in statements[0]
run = db.scalar(select(AgentAssetTestRun).where(AgentAssetTestRun.id == result.id))
assert run is not None
assert run.tenant_id == "tenant-a"
assert run.scope == "tenant"
assert run.input_json["target_tenant_id"] == "tenant-a"
with pytest.raises(LookupError, match="Asset not found"):
service.run_risk_rule_scenario_test(
platform_asset.id,
AgentAssetRiskRuleScenarioTestRequest(
target_tenant_id="tenant-b",
intent="最近 30 天",
),
actor="employee:employee-reviewer-a",
)
def test_release_reviewer_identity_is_stable_when_display_name_changes() -> None:
first = CurrentUserContext(
username="reviewer",
name="显示名一",
role_codes=["manager"],
is_admin=False,
tenant_id="tenant-a",
employee_id="employee-reviewer",
)
renamed = CurrentUserContext(
username="reviewer-renamed",
name="显示名二",
role_codes=["manager"],
is_admin=False,
tenant_id="tenant-a",
employee_id="employee-reviewer",
)
assert stable_user_principal(first) == "employee:employee-reviewer"
assert stable_user_principal(first) == stable_user_principal(renamed)
assert release_actor(first) == release_actor(renamed)
def test_onlyoffice_callback_is_tenant_bound_and_one_time(
monkeypatch: pytest.MonkeyPatch,
) -> None:
factory = _factory()
jwt_secret = "agent-asset-onlyoffice-test-secret-32-bytes"
with factory() as db:
assets = _seed(db)
tenant_asset_id = assets["tenant-a"].id
other_asset_id = assets["tenant-b"].id
tokens = AgentAssetOnlyOfficeSessionService(
db,
jwt_secret=jwt_secret,
).issue(
tenant_id="tenant-a",
resource_scope="tenant",
asset_id=tenant_asset_id,
document_key="tenant-a-document-key",
document_version="v1.0.0",
document_fingerprint="tenant-a-fingerprint",
writable=True,
actor="employee:onlyoffice-editor",
)
monkeypatch.setattr(
"app.services.agent_asset_onlyoffice.resolve_onlyoffice_settings",
lambda: OnlyOfficeRuntimeConfig(
enabled=True,
public_url="https://onlyoffice.example.com",
backend_url="https://backend.example.com",
jwt_secret=jwt_secret,
),
)
def _save_scoped_callback(self, *, claimed, download_url):
assert download_url == "https://onlyoffice.example.com/download/tenant-a.xlsx"
assert claimed.tenant_id == "tenant-a"
assert self.repository.get(tenant_asset_id) is not None
assert self.repository.get(other_asset_id) is None
monkeypatch.setattr(
AgentAssetService,
"_save_current_rule_spreadsheet_callback",
_save_scoped_callback,
)
client = _client(factory)
payload = {
"status": 2,
"url": "https://onlyoffice.example.com/download/tenant-a.xlsx",
"key": "tenant-a-document-key",
}
missing_token = client.post(
f"/agent-assets/{tenant_asset_id}/spreadsheet/onlyoffice/callback",
json=payload,
)
assert missing_token.status_code == 422
cross_asset = client.post(
f"/agent-assets/{other_asset_id}/spreadsheet/onlyoffice/callback",
params={"access_token": tokens.callback_token},
json=payload,
)
assert cross_asset.status_code == 401
first = client.post(
f"/agent-assets/{tenant_asset_id}/spreadsheet/onlyoffice/callback",
params={"access_token": tokens.callback_token},
json=payload,
)
assert first.status_code == 200
replay = client.post(
f"/agent-assets/{tenant_asset_id}/spreadsheet/onlyoffice/callback",
params={"access_token": tokens.callback_token},
json=payload,
)
assert replay.status_code == 409
def test_tenant_security_migration_has_fixed_ancestry_and_fail_closed_downgrade() -> None:
migration = Path(
"/app/server/alembic/versions/20260717_0026_agent_asset_tenant_security.py"
).read_text(encoding="utf-8")
assert 'revision: str = "20260717_0026"' in migration
assert 'down_revision: str | None = "20260717_0025"' in migration
assert "uq_agent_assets_tenant_scope_code" in migration
assert "agent_asset_onlyoffice_sessions" in migration
assert "target_tenant_id" in migration
assert "_require_lossless_downgrade()" in migration
assert "_require_no_onlyoffice_sessions()" in migration
def _claim(tenant_id: str, claim_no: str) -> ExpenseClaim:
return ExpenseClaim(
tenant_id=tenant_id,
claim_no=claim_no,
employee_name="测试员工",
department_name="测试部门",
expense_type="差旅费",
reason="真实场景测试",
location="北京",
amount=100,
currency="CNY",
invoice_count=0,
occurred_at=datetime.now(UTC),
status="submitted",
)