feat: add AI research chatbot with confirmed business tools

This commit is contained in:
yuxuanhui
2026-09-07 23:02:55 +08:00
parent 3cd280d068
commit 79ab20b4ea
44 changed files with 4647 additions and 197 deletions
+54
View File
@@ -0,0 +1,54 @@
"""Deterministic model for isolated acceptance; never calls a provider or platform."""
import asyncio
import json
from contextlib import asynccontextmanager
from uuid import uuid4
from pydantic_ai.messages import ToolReturnPart, UserPromptPart
from pydantic_ai.models.function import DeltaToolCall, FunctionModel
async def fake_stream(messages, info):
latest = max(
(i for i, m in enumerate(messages) if any(isinstance(p, UserPromptPart) for p in m.parts)), default=0
)
text = " ".join(
str(p.content) for m in messages[latest:] for p in m.parts if isinstance(p, UserPromptPart)
)
returns = [p for m in messages[latest:] for p in m.parts if isinstance(p, ToolReturnPart)]
if returns and "LOOP" not in text:
if returns[-1].tool_name == "capability_probe":
yield str(returns[-1].content)
else:
yield "操作结果已返回。"
yield "请查看下方业务记录与数据来源。"
return
if any(t.name == "capability_probe" for t in info.function_tools):
name, args = "capability_probe", {}
elif "READY" in text:
yield "REA"
yield "DY"
return
elif "SLOW" in text:
yield "正在查询"
await asyncio.sleep(2)
yield ",查询完成。"
return
elif "批量" in text:
name, args = "bulk_update_research", {"alpha_ids": ["a0000", "a0001"], "add_tags": ["AI"]}
elif "修改" in text or "update" in text:
alpha_id = "TEST0001" if "TEST0001" in text else "a0000"
name, args = "update_research", {"alpha_id": alpha_id, "changes": {"note": "AI 测试研究记录"}}
elif "同步" in text:
name, args = "create_sync_job", {"kind": "full_sync"}
elif "PnL" in text:
name, args = "get_alpha_pnl", {"alpha_id": "TEST0001" if "TEST" in text else "a0000"}
else:
name, args = "search_alphas", {"filters": {"turnover_max": 0.15, "limit": 5}}
yield {0: DeltaToolCall(name=name, json_args=json.dumps(args), tool_call_id=uuid4().hex)}
@asynccontextmanager
async def fake_model(config, settings):
yield FunctionModel(stream_function=fake_stream, model_name="test-model")
+4 -1
View File
@@ -11,6 +11,7 @@ from app.config import Settings
from app.main import create_app
from app.models import Base
from app.worldquant import WqClient
from tests.ai_fake import fake_model
TEST_PASSWORD = "browser-test-password"
@@ -120,7 +121,9 @@ def create_test_app():
record = next((r for r in records if path == f"/alphas/{r['id']}"), None)
return httpx.Response(200, json=record) if record else httpx.Response(404)
application = create_app(settings, WqClient(settings, transport=httpx.MockTransport(upstream)))
application = create_app(
settings, WqClient(settings, transport=httpx.MockTransport(upstream)), ai_model_factory=fake_model
)
original_lifespan = application.router.lifespan_context
@asynccontextmanager
+123 -1
View File
@@ -89,6 +89,7 @@ asyncio.run(seed())
"PATCH",
{
"note": "persistent local note",
"version": 1,
"tags": ["docker-verified"],
"state": "candidate",
"favorite": True,
@@ -105,6 +106,78 @@ asyncio.run(seed())
assert detail["sharpe"] == 3 and detail["research"]["state"] == "candidate"
assert detail["research"]["tags"] == ["docker-verified"]
print("PASS: container replacement preserves session and data; snapshot update preserves research")
# This synthetic service runs inside the disposable backend container. No real provider is used.
model_server = Path(__file__).with_name("model_protocol.py").read_text()
def start_model():
run(["exec", "-d", "-T", "backend", "python", "-c", model_server])
run(
[
"exec",
"-T",
"backend",
"python",
"-c",
"import socket,time\nfor _ in range(40):\n try:\n socket.create_connection(('127.0.0.1',19010),timeout=1).close();break\n except OSError: time.sleep(.1)\nelse: raise RuntimeError('Mock model did not start')",
]
)
start_model()
for protocol in ("chat_completions", "responses"):
config = {"base_url": "http://127.0.0.1:19010/v1", "model": "mock-model", "protocol": protocol}
json.load(request("/api/v1/ai/settings", "PUT", config | {"api_key": "synthetic-model-key"}))
tested = json.load(request("/api/v1/ai/settings/test", "POST"))
assert tested["ready"], tested
json.load(request("/api/v1/ai/settings", "PUT", config | {"enabled": True}))
conversation = json.load(request("/api/v1/ai/conversations", "POST"))["id"]
started = time.monotonic()
stream = request(
f"/api/v1/ai/conversations/{conversation}/runs", "POST", {"request_id": "stream", "message": "SLOW"}
)
run_id = stream.headers["X-AI-Run-ID"]
assert stream.headers["x-vercel-ai-ui-message-stream"] == "v1"
while b'"text-delta"' not in stream.readline():
assert time.monotonic() - started < 8
assert json.load(request(f"/api/v1/ai/runs/{run_id}"))["status"] == "running"
stream.close()
for _ in range(40):
if json.load(request(f"/api/v1/ai/runs/{run_id}"))["status"] == "completed":
break
time.sleep(0.25)
else:
raise AssertionError("Disconnected execution did not complete")
stream = request(
f"/api/v1/ai/conversations/{conversation}/runs",
"POST",
{"request_id": "write", "message": "修改研究记录"},
)
run_id = stream.headers["X-AI-Run-ID"]
stream.read()
pending = json.load(request(f"/api/v1/ai/runs/{run_id}"))
assert pending["status"] == "waiting_approval", pending
assert (
json.load(request("/api/v1/alphas/DOCKER_ACCEPTANCE"))["research"]["note"] == "persistent local note"
)
run(["up", "-d", "--force-recreate", "--wait"])
start_model()
assert json.load(request("/api/v1/ai/settings"))["ready"]
approval = pending["tools"][0]["id"]
for _ in range(2):
request(f"/api/v1/ai/approvals/{approval}/decision", "POST", {"approved": True}).read()
saved = json.load(request("/api/v1/alphas/DOCKER_ACCEPTANCE"))["research"]
assert saved["note"] == "AI verified note" and saved["version"] == 3
try:
request("/api/v1/alphas/DOCKER_ACCEPTANCE/research", "PATCH", {"version": 2, "note": "stale"})
raise AssertionError("A stale edit overwrote AI research")
except urllib.error.HTTPError as error:
assert error.code == 409
print(
"PASS: both real SDK protocols over mock HTTP; Caddy streams before completion; disconnect recovery"
)
print(
"PASS: pending approval survives container replacement; duplicate confirmation writes once; PostgreSQL version conflict"
)
dump = run(["exec", "-T", "db", "pg_dump", "-U", "wq", "-d", "wq", "-Fc", "--no-owner"])
backup = Path(".local/docker-acceptance.dump")
fd = os.open(backup, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
@@ -147,8 +220,57 @@ asyncio.run(seed())
.decode()
.strip()
)
assert rows == "persistent local note|candidate"
assert rows == "AI verified note|candidate"
ai_rows = (
run(
[
"exec",
"-T",
"db",
"psql",
"-U",
"wq",
"-d",
"wq_acceptance_restore",
"-At",
"-c",
"SELECT (SELECT count(*) FROM ai_conversations), (SELECT count(*) FROM ai_runs), (SELECT count(*) FROM ai_tool_calls), (SELECT count(*) FROM ai_settings WHERE api_key_encrypted IS NOT NULL)",
]
)
.decode()
.strip()
)
assert ai_rows == "1|2|1|1", ai_rows
print("PASS: PostgreSQL custom-format backup restores records into independent database")
# Downgrade only the scratch restore database, then verify upgrading existing 0001 research.
migration_env = (
"DATABASE_URL=postgresql+asyncpg://wq:"
+ values["POSTGRES_PASSWORD"]
+ "@db:5432/wq_acceptance_restore"
)
for command in (["downgrade", "0001"], ["upgrade", "head"], ["check"]):
run(["exec", "-T", "-e", migration_env, "backend", "alembic", *command])
versioned = (
run(
[
"exec",
"-T",
"db",
"psql",
"-U",
"wq",
"-d",
"wq_acceptance_restore",
"-At",
"-c",
"SELECT note || '|' || version FROM research WHERE alpha_id='DOCKER_ACCEPTANCE'",
]
)
.decode()
.strip()
)
assert versioned == "AI verified note|1", versioned
print("PASS: existing 0001 research upgrades with version 1 and unchanged content; Alembic model parity")
config = json.loads(run(["-f", "compose.public.yaml", "config", "--format", "json"]))
assert config["services"]["backend"]["environment"]["COOKIE_SECURE"] == "true"
assert config["services"]["backend"]["environment"]["PUBLIC_ORIGIN"].startswith("https://")
+150
View File
@@ -0,0 +1,150 @@
"""Synthetic OpenAI wire protocol server; no outbound network or business access."""
import json
def model_events(body, protocol):
"""Return concrete SSE frames used by the real SDK adapters in acceptance tests."""
inputs = body.get("messages", body.get("input", []))
outputs = [p for p in inputs if p.get("role") == "tool" or p.get("type") == "function_call_output"]
prompt = json.dumps(inputs, ensure_ascii=False)
tool = None
text = "READY"
if outputs:
text = outputs[-1].get("content", outputs[-1].get("output", ""))
elif "capability_probe" in json.dumps(body.get("tools", [])):
tool = ("capability_probe", {})
elif "修改" in prompt:
tool = ("update_research", {"alpha_id": "DOCKER_ACCEPTANCE", "changes": {"note": "AI verified note"}})
elif "查询" in prompt:
tool = ("search_alphas", {"filters": {"limit": 5, "turnover_max": 0.15}})
elif "SLOW" in prompt:
text = "STREAM READY"
if protocol == "chat_completions":
base = {
"id": "chat-mock",
"object": "chat.completion.chunk",
"created": 1788739200,
"model": body["model"],
}
delta = (
{
"tool_calls": [
{
"index": 0,
"id": "call-mock",
"type": "function",
"function": {"name": tool[0], "arguments": json.dumps(tool[1])},
}
]
}
if tool
else {"content": text}
)
frames = [
base | {"choices": [{"index": 0, "delta": delta, "finish_reason": None}]},
base
| {
"choices": [{"index": 0, "delta": {}, "finish_reason": "tool_calls" if tool else "stop"}],
"usage": {"prompt_tokens": 12, "completion_tokens": 6, "total_tokens": 18},
},
]
return ["data: " + json.dumps(frame) + "\n\n" for frame in frames] + ["data: [DONE]\n\n"]
response = {
"id": "resp-mock",
"object": "response",
"created_at": 1788739200,
"model": body["model"],
"status": "in_progress",
"output": [],
"error": None,
"incomplete_details": None,
"usage": None,
}
frames = [{"type": "response.created", "response": response}]
if tool:
item = {
"id": "fc-mock",
"type": "function_call",
"call_id": "call-mock",
"name": tool[0],
"arguments": "",
"status": "in_progress",
}
frames += [
{"type": "response.output_item.added", "output_index": 0, "item": item},
{
"type": "response.function_call_arguments.delta",
"item_id": "fc-mock",
"output_index": 0,
"delta": json.dumps(tool[1]),
},
{
"type": "response.output_item.done",
"output_index": 0,
"item": item | {"arguments": json.dumps(tool[1]), "status": "completed"},
},
]
else:
item = {
"id": "msg-mock",
"type": "message",
"role": "assistant",
"status": "in_progress",
"content": [],
}
frames += [
{"type": "response.output_item.added", "output_index": 0, "item": item},
{
"type": "response.output_text.delta",
"item_id": "msg-mock",
"output_index": 0,
"content_index": 0,
"delta": text,
"logprobs": [],
},
]
frames.append(
{
"type": "response.completed",
"response": response
| {
"status": "completed",
"usage": {
"input_tokens": 12,
"output_tokens": 6,
"total_tokens": 18,
"input_tokens_details": {"cached_tokens": 0},
"output_tokens_details": {"reasoning_tokens": 0},
},
},
}
)
return [
"event: " + frame["type"] + "\ndata: " + json.dumps(frame | {"sequence_number": i}) + "\n\n"
for i, frame in enumerate(frames)
]
if __name__ == "__main__":
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
class Handler(BaseHTTPRequestHandler):
def do_POST(self):
body = json.loads(self.rfile.read(int(self.headers["Content-Length"])))
self.send_response(200)
self.send_header("Content-Type", "text/event-stream")
self.end_headers()
for frame in model_events(
body, "responses" if self.path.endswith("/responses") else "chat_completions"
):
self.wfile.write(frame.encode())
self.wfile.flush()
if "SLOW" in json.dumps(body):
time.sleep(1)
def log_message(self, *args):
pass
ThreadingHTTPServer(("127.0.0.1", 19010), Handler).serve_forever()
+331
View File
@@ -0,0 +1,331 @@
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
+76
View File
@@ -0,0 +1,76 @@
import json
import httpx
import pytest
from app.ai.provider import model_connection
from app.ai.provider import test_capabilities as check_capabilities
from app.models import AISettings
from app.security import cipher
from tests.model_protocol import model_events
@pytest.mark.parametrize("protocol", ["chat_completions", "responses"])
async def test_actual_provider_protocol_and_tool_roundtrip(app, protocol):
paths = []
def gateway(request):
paths.append(request.url.path)
assert request.headers["authorization"] == "Bearer synthetic-key"
body = json.loads(request.content)
assert body["model"] == "mock-model" and body["stream"] is True
return httpx.Response(
200, headers={"content-type": "text/event-stream"}, content="".join(model_events(body, protocol))
)
config = AISettings(
base_url="http://model.test/v1",
model="mock-model",
protocol=protocol,
api_key_encrypted=cipher(app.state.settings).encrypt(b"synthetic-key").decode(),
)
async with model_connection(config, app.state.settings, httpx.MockTransport(gateway)) as model:
result = await check_capabilities(model)
assert all(item["ok"] for item in result.values()), result
assert paths == ["/v1/" + ("chat/completions" if protocol == "chat_completions" else "responses")] * 3
@pytest.mark.parametrize("protocol", ["chat_completions", "responses"])
@pytest.mark.parametrize("failure", [401, 404, 429, "timeout", "broken", "truncated", "no-tools"])
async def test_provider_failures_are_safe(app, protocol, failure, caplog):
def gateway(request):
if isinstance(failure, int):
return httpx.Response(failure, json={"error": {"message": "synthetic-key private provider body"}})
if failure == "timeout":
raise httpx.ReadTimeout("synthetic-key private provider body")
if failure == "broken":
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
content="data: invalid-private-synthetic-key\n\n",
)
body = json.loads(request.content)
body["tools"] = []
if failure == "truncated":
frames = model_events(body, protocol)
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
content="".join(frames[:-2] if protocol == "chat_completions" else frames[:-1]),
)
return httpx.Response(
200, headers={"content-type": "text/event-stream"}, content="".join(model_events(body, protocol))
)
config = AISettings(
base_url="http://model.test/v1",
model="mock-model",
protocol=protocol,
api_key_encrypted=cipher(app.state.settings).encrypt(b"synthetic-key").decode(),
)
async with model_connection(config, app.state.settings, httpx.MockTransport(gateway)) as model:
result = await check_capabilities(model)
assert not result["tools"]["ok"], result
assert "synthetic-key" not in json.dumps(result) + caplog.text
if failure == "no-tools":
assert result["answer"]["ok"] and result["stream"]["ok"]
+4 -3
View File
@@ -113,7 +113,7 @@ async def test_local_research_survives_snapshot_and_bulk_is_atomic(app, logged_i
await seed(app, 3)
r = await logged_in.patch(
f"{PREFIX}/alphas/a0000/research",
json={"note": "keep hypothesis", "tags": ["a", "a"], "state": "candidate", "favorite": True},
json={"version": 1, "note": "keep hypothesis", "tags": ["a", "a"], "state": "candidate", "favorite": True},
)
assert r.status_code == 200
async with app.state.sessions() as db:
@@ -124,7 +124,7 @@ async def test_local_research_survives_snapshot_and_bulk_is_atomic(app, logged_i
assert detail["research"]["note"] == "keep hypothesis" and detail["research"]["favorite"]
assert detail["research"]["tags"] == ["a"] and detail["research"]["state"] == "candidate"
invalid = await logged_in.patch(
f"{PREFIX}/alphas/research/bulk", json={"alpha_ids": ["a0000", "missing"], "add_tags": ["bad"]}
f"{PREFIX}/alphas/research/bulk", json={"alpha_ids": ["a0000", "missing"], "add_tags": ["bad"], "versions": {"a0000": 2, "missing": 1}}
)
assert invalid.status_code == 404
assert (await logged_in.get(f"{PREFIX}/alphas?tag=bad")).json()["total"] == 0
@@ -132,6 +132,7 @@ async def test_local_research_survives_snapshot_and_bulk_is_atomic(app, logged_i
f"{PREFIX}/alphas/research/bulk",
json={
"alpha_ids": ["a0000", "a0001"],
"versions": {"a0000": 2, "a0001": 1},
"add_tags": ["new"],
"remove_tags": ["a"],
"state": "optimizing",
@@ -141,7 +142,7 @@ async def test_local_research_survives_snapshot_and_bulk_is_atomic(app, logged_i
result = (await logged_in.get(f"{PREFIX}/alphas?tag=new&research_state=optimizing")).json()
assert result["total"] == 2
assert (await logged_in.get(f"{PREFIX}/alphas?tag=ne")).json()["total"] == 0
await logged_in.patch(f"{PREFIX}/alphas/a0000/research", json={"note": "updated"})
await logged_in.patch(f"{PREFIX}/alphas/a0000/research", json={"note": "updated", "version": 3})
detail = (await logged_in.get(f"{PREFIX}/alphas/a0000")).json()
assert detail["research"]["favorite"] and detail["research"]["tags"] == ["new"]