206 lines
8.3 KiB
Python
206 lines
8.3 KiB
Python
|
|
"""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, TemplateInput
|
||
|
|
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 SCOPE
|
||
|
|
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, "准备输入后的校验失败")
|
||
|
|
|
||
|
|
monkeypatch.setattr(ResearchBuilder, "input_page", unavailable_page)
|
||
|
|
app.state.ai.model_factory = single_tool_factory(
|
||
|
|
"prepare_research_input",
|
||
|
|
{
|
||
|
|
"scope": SCOPE,
|
||
|
|
"dataset_id": "TEST_FIN",
|
||
|
|
"collection_version": fixed_input["collection_version"],
|
||
|
|
"field_ids": ["TEST_FIN_001"],
|
||
|
|
},
|
||
|
|
)
|
||
|
|
_, 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(TemplateInput)) == 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
|