feat: add AI research chatbot with confirmed business tools
This commit is contained in:
@@ -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")
|
||||
@@ -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
|
||||
|
||||
@@ -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://")
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user