refactor: unify AI capabilities and workspace integration
This commit is contained in:
@@ -0,0 +1,205 @@
|
||||
"""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
|
||||
@@ -268,8 +268,12 @@ async def test_dynamic_platform_scopes_and_validation(catalog):
|
||||
assert (await sync(catalog, scope={**SCOPE, "region": "IND", "universe": "TOP500"}))["status"] == "completed"
|
||||
invalid = await client.post(BASE + "/sync-jobs", json={"scope": {**SCOPE, "region": "IND"}})
|
||||
assert invalid.status_code == 422
|
||||
from app.ai.tools import EmptyArgs, read_tool
|
||||
ai = await read_tool(None, "get_catalog_scopes", EmptyArgs(), runner.client)
|
||||
from app.ai.capabilities import ToolContext
|
||||
from app.ai.tools import CAPABILITIES
|
||||
from app.business import Business
|
||||
|
||||
async with runner.sessions() as db:
|
||||
ai = await CAPABILITIES["get_catalog_scopes"].invoke(ToolContext(Business(db), runner.client), {})
|
||||
assert ai["instrument_options"] == options["instrument_options"]
|
||||
assert ai["_meta"]["source"] == "worldquant_platform"
|
||||
runner.client.disconnect()
|
||||
|
||||
@@ -6,7 +6,8 @@ import pytest
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from app.ai.tools import CATALOG, read_tool
|
||||
from app.ai.capabilities import ToolContext
|
||||
from app.ai.tools import CAPABILITIES
|
||||
from app.alphas import upsert_alpha
|
||||
from app.backtests.contracts import PreviewInput, RerunInput, SubsetInput
|
||||
from app.business import Business
|
||||
@@ -137,7 +138,7 @@ async def test_invalid_construction_has_no_partial_preview(logged_in, app, fixed
|
||||
async def test_input_pagination_old_types_and_explicit_exclusions(app, logged_in, catalog, fixed_input):
|
||||
async def tool(name, args):
|
||||
async with app.state.sessions.begin() as db:
|
||||
return await read_tool(Business(db), name, CATALOG[name][0].model_validate(args))
|
||||
return await CAPABILITIES[name].invoke(ToolContext(Business(db)), args)
|
||||
|
||||
page = await tool("get_research_input", {"input_id": fixed_input["id"], "offset": 100, "limit": 25})
|
||||
assert page["field_count"] == page["total"] == 123 and len(page["items"]) == 23
|
||||
|
||||
Reference in New Issue
Block a user