Files

203 lines
8.4 KiB
Python
Raw Permalink Normal View History

"""Capability policy through the real executor, database and persisted UI stream."""
import json
from dataclasses import replace
import pytest
from fastapi import HTTPException
from sqlalchemy import func, select
from app.ai.capabilities import ToolContext, assemble
from app.ai.tools import CAPABILITIES
from app.business import Business
from app.models import AIMessage, AIRun, AIToolCall, Alpha, Research, ResearchInputSnapshot
from app.research.service import ResearchBuilder
from tests.test_ai import configure, single_tool_factory, start
from tests.test_api import seed
from tests.test_catalog import catalog as catalog_fixture
from tests.test_research_integration import fixed_input as fixed_input_fixture
catalog = catalog_fixture
fixed_input = fixed_input_fixture
@pytest.mark.parametrize(
"changes",
[
{"effect": "confirm"},
{"effect": "unclassified"},
{"refresh": ("alphas",)},
{"after_commit": lambda runner, result: None},
{"renderer": ""},
],
)
def test_incomplete_or_ambiguous_policy_fails_at_assembly(changes):
with pytest.raises(ValueError):
replace(CAPABILITIES["get_alpha"], **changes)
def test_duplicate_names_cannot_replace_an_existing_capability():
capability = CAPABILITIES["get_alpha"]
with pytest.raises(ValueError, match="Duplicate capability"):
assemble([[capability], [capability]])
def test_unknown_refresh_target_fails_instead_of_silently_leaving_stale_data():
with pytest.raises(ValueError, match="workspace resource"):
replace(CAPABILITIES["update_research"], refresh=("unknown",))
async def test_confirmation_cannot_be_bypassed_through_invoke(app):
async with app.state.sessions.begin() as db:
with pytest.raises(HTTPException) as exc:
await CAPABILITIES["update_research"].invoke(ToolContext(Business(db)), {})
assert exc.value.status_code == 409
async def test_prepare_failure_rolls_back_artifact_but_keeps_audit(app, logged_in, fixed_input, monkeypatch):
await configure(app, logged_in)
async def unavailable_page(self, *args, **kwargs):
# select_input has already persisted the new input before requesting its result page.
raise HTTPException(422, "准备输入后的校验失败")
updated = await logged_in.patch(f"/api/v1/data-preparations/{fixed_input['preparation_id']}",
json={"version": fixed_input["preparation_version"], "name": "new version"})
assert updated.status_code == 200
monkeypatch.setattr(ResearchBuilder, "input_page", unavailable_page)
app.state.ai.model_factory = single_tool_factory(
"prepare_research_input",
{"items": [{"id": fixed_input["preparation_id"], "version": updated.json()["version"]}]},
)
_, run, _ = await start(app, logged_in, "保存研究输入")
call = run["tools"][0]
assert call["status"] == "failed" and run["status"] == "completed"
assert call["presentation"]["effect"] == "prepare"
assert call["result"]["error"] == "准备输入后的校验失败"
async with app.state.sessions() as db:
assert await db.scalar(select(func.count()).select_from(ResearchInputSnapshot)) == 1
assert (await db.get(AIToolCall, call["id"])).status == "failed"
async def test_model_summary_does_not_truncate_persisted_card(app, logged_in):
await configure(app, logged_in)
await seed(app, 1)
expression = " + ".join(["rank(close)"] * 300)
async with app.state.sessions.begin() as db:
(await db.get(Alpha, "a0000")).expression = expression
app.state.ai.model_factory = single_tool_factory("get_alpha", {"alpha_id": "a0000"})
conversation, run, stream = await start(app, logged_in)
call = run["tools"][0]
assert call["result"]["expression"] == expression
assert call["presentation"]["refresh"] == []
assert call["presentation"]["label"] == "读取 Alpha"
assert '"presentation"' in stream.text
async with app.state.sessions() as db:
saved = await db.get(AIRun, run["id"])
returns = [p for m in saved.model_messages for p in m["parts"] if p["part_kind"] == "tool-return"]
assert returns[0]["content"]["_meta"]["truncated"] is True
assert len(returns[0]["content"]["expression"]) < len(expression)
history = (await logged_in.get(f"/api/v1/ai/conversations/{conversation}")).json()
card = next(p["data"] for m in history["messages"] for p in m["parts"] if p["type"] == "data-tool")
assert card["result"]["expression"] == expression
async def test_historical_approval_hydrates_presentation_and_executes_once(app, logged_in):
await configure(app, logged_in)
await seed(app, 1)
conversation, run, _ = await start(app, logged_in, "修改")
async with app.state.sessions.begin() as db:
for message in await db.scalars(select(AIMessage).where(AIMessage.run_id == run["id"])):
parts = json.loads(json.dumps(message.parts))
for part in parts:
if part["type"] == "data-tool":
part["data"].pop("presentation", None)
message.parts = parts
await app.state.ai.start()
history = (await logged_in.get(f"/api/v1/ai/conversations/{conversation}")).json()
call = history["runs"][0]["tools"][0]
assert call["presentation"]["effect"] == "confirm"
assert call["presentation"]["refresh"] == ["alphas"]
for _ in range(2):
response = await logged_in.post(
f"/api/v1/ai/approvals/{call['id']}/decision", json={"approved": True}
)
assert response.status_code == 200
async with app.state.sessions() as db:
assert (await db.get(Research, "a0000")).version == 2
async def test_removed_capability_cannot_execute_a_historical_approval(app, logged_in, monkeypatch):
await configure(app, logged_in)
await seed(app, 1)
_, run, _ = await start(app, logged_in, "修改")
call = run["tools"][0]
monkeypatch.delitem(CAPABILITIES, "update_research")
response = await logged_in.post(f"/api/v1/ai/approvals/{call['id']}/decision", json={"approved": True})
assert response.status_code == 200
snapshot = (await logged_in.get(f"/api/v1/ai/runs/{run['id']}")).json()
assert snapshot["tools"][0]["status"] == "failed"
assert snapshot["tools"][0]["presentation"]["effect"] == "unavailable"
async with app.state.sessions() as db:
assert (await db.get(Research, "a0000")).version == 1
async def test_notification_observes_committed_operation_and_audit(app, logged_in, monkeypatch):
await configure(app, logged_in)
await seed(app, 1)
notified = []
async def after_commit(runner, result):
async with app.state.sessions() as db:
assert (await db.get(Research, "a0000")).version == 2
call = await db.scalar(select(AIToolCall).where(AIToolCall.name == "update_research"))
assert call.status == "completed" and call.result == result
notified.append(result)
monkeypatch.setitem(
CAPABILITIES,
"update_research",
replace(
CAPABILITIES["update_research"],
after_commit=after_commit,
),
)
_, run, _ = await start(app, logged_in, "修改")
for _ in range(2):
await logged_in.post(
f"/api/v1/ai/approvals/{run['tools'][0]['id']}/decision", json={"approved": True}
)
assert len(notified) == 1
async def test_failed_notification_keeps_commit_and_finishes_chat(app, logged_in, monkeypatch):
await configure(app, logged_in)
await seed(app, 1)
notified = []
async def unavailable(runner, result):
notified.append(result)
raise RuntimeError("synthetic internal notification failure")
monkeypatch.setitem(
CAPABILITIES,
"update_research",
replace(CAPABILITIES["update_research"], after_commit=unavailable),
)
_, run, _ = await start(app, logged_in, "修改")
for _ in range(2):
response = await logged_in.post(
f"/api/v1/ai/approvals/{run['tools'][0]['id']}/decision",
json={"approved": True},
)
assert response.status_code == 200
snapshot = (await logged_in.get(f"/api/v1/ai/runs/{run['id']}")).json()
assert snapshot["status"] == "completed"
call = snapshot["tools"][0]
assert call["status"] == "completed"
assert "操作已保存" in call["result"]["_warning"]
assert "synthetic internal" not in json.dumps(snapshot)
assert len(notified) == 1
async with app.state.sessions() as db:
assert (await db.get(Research, "a0000")).version == 2