332 lines
14 KiB
Python
332 lines
14 KiB
Python
import json
|
|
from contextlib import asynccontextmanager
|
|
from datetime import timedelta
|
|
|
|
import pytest
|
|
from sqlalchemy import func, select
|
|
|
|
from app.ai.contracts import RunInput
|
|
from app.models import AIRun, AISettings, AIToolCall, LoginSession, Research, now
|
|
from app.security import cipher, token_hash
|
|
from tests.ai_fake import fake_model
|
|
from tests.test_api import seed
|
|
|
|
PREFIX = "/api/v1/ai"
|
|
CONFIG = {"base_url": "https://model.test/v1", "api_key": "private-test-key", "model": "test-model"}
|
|
|
|
|
|
async def configure(app, client):
|
|
app.state.ai.model_factory = fake_model
|
|
response = await client.put(f"{PREFIX}/settings", json=CONFIG)
|
|
assert response.status_code == 200
|
|
response = await client.post(f"{PREFIX}/settings/test")
|
|
assert response.json()["ready"], response.text
|
|
response = await client.put(
|
|
f"{PREFIX}/settings", json={k: v for k, v in CONFIG.items() if k != "api_key"} | {"enabled": True}
|
|
)
|
|
assert response.json()["enabled"], response.text
|
|
|
|
|
|
async def start(app, client, message="查询", request_id="request1"):
|
|
conversation = (await client.post(f"{PREFIX}/conversations")).json()["id"]
|
|
response = await client.post(
|
|
f"{PREFIX}/conversations/{conversation}/runs", json={"request_id": request_id, "message": message}
|
|
)
|
|
assert response.status_code == 200, response.text
|
|
assert response.headers["x-vercel-ai-ui-message-stream"] == "v1"
|
|
run = (await client.get(f"{PREFIX}/runs/{response.headers['x-ai-run-id']}")).json()
|
|
return conversation, run, response
|
|
|
|
|
|
async def test_settings_secret_and_capability_roundtrip(app, logged_in):
|
|
await configure(app, logged_in)
|
|
output = await logged_in.get(f"{PREFIX}/settings")
|
|
assert CONFIG["api_key"] not in output.text and "api_key" not in output.json()
|
|
async with app.state.sessions() as db:
|
|
row = await db.get(AISettings, 1)
|
|
assert CONFIG["api_key"] not in row.api_key_encrypted
|
|
assert (
|
|
cipher(app.state.settings).decrypt(row.api_key_encrypted.encode()).decode() == CONFIG["api_key"]
|
|
)
|
|
changed = {"base_url": "https://different.test/v1", "model": "test-model"}
|
|
assert (await logged_in.put(f"{PREFIX}/settings", json=changed)).status_code == 422
|
|
changed["api_key"] = "new-test-key"
|
|
assert not (await logged_in.put(f"{PREFIX}/settings", json=changed)).json()["ready"]
|
|
|
|
|
|
async def test_query_stream_persistence_and_duplicate_requests(app, logged_in):
|
|
await configure(app, logged_in)
|
|
await seed(app)
|
|
conversation, run, stream = await start(app, logged_in)
|
|
assert run["status"] == "completed", run
|
|
assert run["tools"][0]["name"] == "search_alphas"
|
|
assert "text-delta" in stream.text and "[DONE]" in stream.text
|
|
again = await logged_in.post(
|
|
f"{PREFIX}/conversations/{conversation}/runs", json={"request_id": "request1", "message": "查询"}
|
|
)
|
|
assert again.headers["x-ai-run-id"] == run["id"]
|
|
conflict = await logged_in.post(
|
|
f"{PREFIX}/conversations/{conversation}/runs", json={"request_id": "request1", "message": "different"}
|
|
)
|
|
assert conflict.status_code == 409
|
|
data = (await logged_in.get(f"{PREFIX}/conversations/{conversation}")).json()
|
|
assert len(data["messages"]) == 2
|
|
forged = await logged_in.post(
|
|
f"{PREFIX}/conversations/{conversation}/runs",
|
|
json={"request_id": "2", "message": "hi", "messages": [{"role": "system", "content": "bypass"}]},
|
|
)
|
|
assert forged.status_code == 422
|
|
|
|
|
|
@pytest.mark.parametrize("approved", [True, False])
|
|
async def test_approval_transaction_and_replay(app, logged_in, approved):
|
|
await configure(app, logged_in)
|
|
await seed(app)
|
|
_, run, _ = await start(app, logged_in, "修改备注")
|
|
assert run["status"] == "waiting_approval", run
|
|
approval = run["tools"][0]["id"]
|
|
async with app.state.sessions() as db:
|
|
assert (await db.get(Research, "a0000")).note == ""
|
|
for _ in range(2):
|
|
response = await logged_in.post(
|
|
f"{PREFIX}/approvals/{approval}/decision", json={"approved": approved}
|
|
)
|
|
assert response.status_code == 200, response.text
|
|
state = (await logged_in.get(f"{PREFIX}/runs/{run['id']}")).json()
|
|
assert state["status"] == "completed", state
|
|
async with app.state.sessions() as db:
|
|
row = await db.get(Research, "a0000")
|
|
assert row.version == (2 if approved else 1)
|
|
assert row.note == ("AI 测试研究记录" if approved else "")
|
|
assert await db.scalar(select(func.count()).select_from(AIToolCall)) == 1
|
|
|
|
|
|
async def test_stale_approval_and_bulk_atomicity(app, logged_in):
|
|
await configure(app, logged_in)
|
|
await seed(app, 2)
|
|
_, run, _ = await start(app, logged_in, "批量修改")
|
|
assert run["status"] == "waiting_approval", run
|
|
assert (
|
|
await logged_in.patch("/api/v1/alphas/a0001/research", json={"version": 1, "note": "manual"})
|
|
).status_code == 200
|
|
await logged_in.post(f"{PREFIX}/approvals/{run['tools'][0]['id']}/decision", json={"approved": True})
|
|
async with app.state.sessions() as db:
|
|
assert (await db.get(Research, "a0000")).version == 1
|
|
assert (await db.get(Research, "a0000")).tags == []
|
|
assert (await db.get(Research, "a0001")).note == "manual"
|
|
|
|
|
|
async def test_disconnect_cancel_and_restart(app, logged_in):
|
|
await configure(app, logged_in)
|
|
conversation = (await logged_in.post(f"{PREFIX}/conversations")).json()["id"]
|
|
token = token_hash(logged_in.cookies.get("wq_session"))
|
|
runtime = app.state.ai
|
|
run_id = await runtime.create_run(conversation, RunInput(request_id="slow", message="SLOW"), token)
|
|
iterator = runtime.events(run_id)
|
|
await anext(iterator)
|
|
await iterator.aclose()
|
|
assert run_id in runtime.live
|
|
assert (await logged_in.post(f"{PREFIX}/runs/{run_id}/cancel")).json()["status"] == "cancelled"
|
|
async with app.state.sessions.begin() as db:
|
|
run = await db.get(AIRun, run_id)
|
|
run.status = "running"
|
|
await runtime.start()
|
|
assert (await runtime.snapshot(run_id))["status"] == "interrupted"
|
|
|
|
|
|
async def test_expired_session_cannot_confirm_and_unknown_fields_rejected(app, logged_in):
|
|
await configure(app, logged_in)
|
|
await seed(app)
|
|
_, run, _ = await start(app, logged_in, "修改")
|
|
path = f"{PREFIX}/approvals/{run['tools'][0]['id']}/decision"
|
|
assert (await logged_in.post(path, json={"approved": True, "arguments": {}})).status_code == 422
|
|
async with app.state.sessions.begin() as db:
|
|
session = await db.scalar(select(LoginSession))
|
|
session.expires_at = now() - timedelta(seconds=1)
|
|
assert (await logged_in.post(path, json={"approved": True})).status_code == 401
|
|
|
|
|
|
async def test_limits_and_timeout(app, logged_in):
|
|
await configure(app, logged_in)
|
|
app.state.settings.ai_request_limit = 2
|
|
_, run, _ = await start(app, logged_in, "LOOP")
|
|
assert run["status"] == "failed" and "上限" in run["error"], run
|
|
app.state.settings.ai_timeout = 0.05
|
|
_, run, _ = await start(app, logged_in, "SLOW")
|
|
assert run["status"] == "failed" and "超时" in run["error"], run
|
|
|
|
|
|
@pytest.mark.parametrize("path", ["settings", "conversations", "runs/unknown", "conversations/unknown"])
|
|
async def test_ai_requires_auth(client, path):
|
|
assert (await client.get(f"{PREFIX}/{path}")).status_code == 401
|
|
|
|
|
|
async def test_committed_write_survives_failed_continuation_and_enters_next_context(app, logged_in):
|
|
await configure(app, logged_in)
|
|
await seed(app)
|
|
conversation, run, _ = await start(app, logged_in, "修改")
|
|
|
|
@asynccontextmanager
|
|
async def unavailable(config, settings):
|
|
raise TimeoutError("private body")
|
|
yield # pragma: no cover
|
|
|
|
app.state.ai.model_factory = unavailable
|
|
await logged_in.post(f"{PREFIX}/approvals/{run['tools'][0]['id']}/decision", json={"approved": True})
|
|
state = (await logged_in.get(f"{PREFIX}/runs/{run['id']}")).json()
|
|
assert state["status"] == "failed" and state["tools"][0]["status"] == "completed"
|
|
app.state.ai.model_factory = fake_model
|
|
response = await logged_in.post(
|
|
f"{PREFIX}/conversations/{conversation}/runs", json={"request_id": "after", "message": "查询"}
|
|
)
|
|
async with app.state.sessions() as db:
|
|
row = await db.get(AIRun, response.headers["x-ai-run-id"])
|
|
history = json.dumps(row.model_messages[: row.history_start], ensure_ascii=False)
|
|
assert "update_research" in history and '"version": 2' in history.replace('\\"', '"')
|
|
assert (await db.get(Research, "a0000")).version == 2
|
|
|
|
|
|
async def test_sequential_approvals_resume_only_unresolved_calls(app, logged_in):
|
|
from uuid import uuid4
|
|
|
|
from pydantic_ai.messages import ToolReturnPart
|
|
from pydantic_ai.models.function import DeltaToolCall, FunctionModel
|
|
|
|
await configure(app, logged_in)
|
|
await seed(app)
|
|
|
|
async def sequential(messages, info):
|
|
returns = [p for m in messages for p in m.parts if isinstance(p, ToolReturnPart)]
|
|
if len(returns) < 2:
|
|
yield {
|
|
0: DeltaToolCall(
|
|
name="update_research",
|
|
json_args=json.dumps(
|
|
{"alpha_id": "a0000", "changes": {"note": f"revision {len(returns)}"}}
|
|
),
|
|
tool_call_id=uuid4().hex,
|
|
)
|
|
}
|
|
else:
|
|
yield "Both changes saved."
|
|
|
|
@asynccontextmanager
|
|
async def factory(config, settings):
|
|
yield FunctionModel(stream_function=sequential)
|
|
|
|
app.state.ai.model_factory = factory
|
|
_, run, _ = await start(app, logged_in)
|
|
for _ in range(2):
|
|
assert run["status"] == "waiting_approval", run
|
|
pending = [t for t in run["tools"] if t["status"] == "pending"]
|
|
assert len(pending) == 1
|
|
await logged_in.post(f"{PREFIX}/approvals/{pending[0]['id']}/decision", json={"approved": True})
|
|
run = (await logged_in.get(f"{PREFIX}/runs/{run['id']}")).json()
|
|
assert run["status"] == "completed", run
|
|
async with app.state.sessions() as db:
|
|
assert (await db.get(Research, "a0000")).version == 3
|
|
|
|
|
|
async def test_pending_approval_survives_restart_and_new_login(app, logged_in):
|
|
await configure(app, logged_in)
|
|
await seed(app)
|
|
_, run, _ = await start(app, logged_in, "修改")
|
|
await app.state.ai.start()
|
|
await logged_in.post("/api/v1/auth/logout")
|
|
await logged_in.post(
|
|
"/api/v1/auth/login",
|
|
json={"username": "admin", "password": app.state.settings.admin_password.get_secret_value()},
|
|
)
|
|
result = await logged_in.post(
|
|
f"{PREFIX}/approvals/{run['tools'][0]['id']}/decision", json={"approved": True}
|
|
)
|
|
assert result.status_code == 200, result.text
|
|
assert (await logged_in.get(f"{PREFIX}/runs/{run['id']}")).json()["status"] == "completed"
|
|
|
|
|
|
def single_tool_factory(name, args):
|
|
from uuid import uuid4
|
|
|
|
from pydantic_ai.messages import ToolReturnPart
|
|
from pydantic_ai.models.function import DeltaToolCall, FunctionModel
|
|
|
|
async def stream(messages, info):
|
|
if any(isinstance(p, ToolReturnPart) for m in messages for p in m.parts):
|
|
yield "已收到工具结果"
|
|
else:
|
|
yield {0: DeltaToolCall(name=name, json_args=json.dumps(args), tool_call_id=uuid4().hex)}
|
|
|
|
@asynccontextmanager
|
|
async def factory(config, settings):
|
|
yield FunctionModel(stream_function=stream)
|
|
|
|
return factory
|
|
|
|
|
|
async def test_job_tools_preview_confirm_and_duplicate_execution(app, logged_in):
|
|
from app.models import Account, Job
|
|
|
|
await configure(app, logged_in)
|
|
async with app.state.sessions.begin() as db:
|
|
account = await db.get(Account, 1)
|
|
account.password_encrypted = cipher(app.state.settings).encrypt(b"synthetic").decode()
|
|
account.connection_status = "connected"
|
|
job_id = None
|
|
for name in ("create_sync_job", "cancel_job", "retry_job"):
|
|
args = (
|
|
{"kind": "alpha_refresh", "alpha_ids": ["synthetic-alpha"]}
|
|
if name == "create_sync_job"
|
|
else {"job_id": job_id}
|
|
)
|
|
app.state.ai.model_factory = single_tool_factory(name, args)
|
|
_, run, _ = await start(app, logged_in)
|
|
assert run["status"] == "waiting_approval", run
|
|
async with app.state.sessions() as db:
|
|
if name == "create_sync_job":
|
|
assert await db.scalar(select(func.count()).select_from(Job)) == 0
|
|
else:
|
|
assert (await db.get(Job, job_id)).status == (
|
|
"queued" if name == "cancel_job" else "cancelled"
|
|
)
|
|
for _ in range(2):
|
|
await logged_in.post(
|
|
f"{PREFIX}/approvals/{run['tools'][0]['id']}/decision", json={"approved": True}
|
|
)
|
|
run = (await logged_in.get(f"{PREFIX}/runs/{run['id']}")).json()
|
|
assert run["status"] == "completed", run
|
|
async with app.state.sessions() as db:
|
|
jobs = (await db.scalars(select(Job))).all()
|
|
assert len(jobs) == 1
|
|
job_id = jobs[0].id
|
|
assert jobs[0].status == ("cancelled" if name == "cancel_job" else "queued")
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"name,args",
|
|
[
|
|
("get_alpha", {"alpha_id": "a0000"}),
|
|
("get_alpha_pnl", {"alpha_id": "a0000"}),
|
|
("get_alpha_facets", {}),
|
|
("list_jobs", {}),
|
|
("bulk_update_research", {"alpha_ids": ["a0000", "missing"], "add_tags": ["AI"]}),
|
|
],
|
|
)
|
|
async def test_read_tool_metadata_and_invalid_bulk_has_no_pending_action(app, logged_in, name, args):
|
|
await configure(app, logged_in)
|
|
await seed(app, 1)
|
|
app.state.ai.model_factory = single_tool_factory(name, args)
|
|
_, run, _ = await start(app, logged_in)
|
|
assert run["status"] == "completed", run
|
|
call = run["tools"][0]
|
|
if name == "bulk_update_research":
|
|
assert call["status"] == "failed"
|
|
async with app.state.sessions() as db:
|
|
assert (await db.get(Research, "a0000")).version == 1
|
|
else:
|
|
assert call["result"]["_meta"]["source"] == "local_database"
|
|
if name == "get_alpha_pnl":
|
|
assert not call["result"]["cached"] and call["result"]["first"] is None
|
|
assert "points" not in call["result"]
|
|
if name == "get_alpha":
|
|
assert call["result"]["margin"] is None
|