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
+39 -156
View File
@@ -11,20 +11,23 @@ from typing import Annotated
from fastapi import APIRouter, Depends, FastAPI, HTTPException, Query, Request, Response
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse, StreamingResponse
from sqlalchemy import delete, func, select, text
from sqlalchemy import delete, select, text
from .alphas import list_statement, sorted_statement, summary
from .ai.routes import router as ai_router
from .ai.runtime import AIRuntime
from .alphas import list_statement, sorted_statement
from .business import Business, notify_job
from .config import Settings
from .db import create_database
from .jobs import ACTIVE, AUTH_KINDS, Runner, create_job
from .models import Account, Admin, Alpha, Job, JobItem, LoginSession, Pnl, Research, ResearchTag, now
from .jobs import AUTH_KINDS, Runner, create_job
from .models import Account, Admin, Job, JobItem, LoginSession
from .schemas import (
AccountOutput,
AlphaDetail,
AlphaFilters,
AlphaPage,
BulkInput,
BulkOutput,
BulkUpdate,
CredentialsInput,
ErrorOutput,
FacetsOutput,
@@ -36,9 +39,8 @@ from .schemas import (
OkOutput,
PnlOutput,
PreferencesInput,
ResearchInput,
ResearchUpdate,
SessionOutput,
normalize_tags,
)
from .security import bootstrap, cipher, issue_session, require_auth, token_hash, valid_password
@@ -60,12 +62,6 @@ def account_output(account):
return {**{k: getattr(account, k) for k in keys}, "configured": bool(account.password_encrypted)}
async def set_tags(db, record, tags):
record.tags = normalize_tags(tags)
await db.execute(delete(ResearchTag).where(ResearchTag.alpha_id == record.alpha_id))
db.add_all(ResearchTag(alpha_id=record.alpha_id, tag=t) for t in record.tags)
def csv_cell(value):
"""Neutralize spreadsheet formulas in untrusted names, expressions, notes and tags."""
if value is None:
@@ -77,18 +73,21 @@ def csv_cell(value):
return value
def create_app(settings=None, wq_client=None):
def create_app(settings=None, wq_client=None, ai_model_factory=None):
settings = settings or Settings()
engine, sessions = create_database(settings.database_url)
runner = Runner(sessions, settings, client=wq_client)
ai_runtime = AIRuntime(sessions, settings, runner, ai_model_factory)
@asynccontextmanager
async def lifespan(app):
async with sessions() as db:
await bootstrap(db, settings)
await ai_runtime.start()
if settings.enable_runner:
await runner.start()
yield
await ai_runtime.stop()
if settings.enable_runner:
await runner.stop()
else:
@@ -103,6 +102,7 @@ def create_app(settings=None, wq_client=None):
)
app.state.engine, app.state.sessions, app.state.runner = engine, sessions, runner
app.state.settings = settings
app.state.ai = ai_runtime
login_failures = defaultdict(list)
@app.exception_handler(RequestValidationError)
@@ -252,45 +252,13 @@ def create_app(settings=None, wq_client=None):
@api.get("/alphas", response_model=AlphaPage, tags=["alphas"])
async def get_alphas(filters: Annotated[AlphaFilters, Query()]):
query = list_statement(filters)
async with sessions() as db:
total = await db.scalar(select(func.count()).select_from(query.subquery()))
rows = (
await db.execute(
sorted_statement(query, filters.sort, filters.direction)
.limit(filters.limit)
.offset(filters.offset)
)
).all()
return {
"items": [summary(a, r) for a, r in rows],
"total": total,
"limit": filters.limit,
"offset": filters.offset,
}
return await Business(db).search_alphas(filters)
@api.get("/alphas/facets", response_model=FacetsOutput, tags=["alphas"])
async def facets():
async with sessions() as db:
result = {}
for key in ("region", "universe", "alpha_type", "language", "status", "stage"):
column = getattr(Alpha, key)
result[key] = list(
(
await db.scalars(
select(column).where(column.is_not(None)).distinct().order_by(column)
)
).all()
)
result["tags"] = list(
(await db.scalars(select(ResearchTag.tag).distinct().order_by(ResearchTag.tag))).all()
)
result["total"] = await db.scalar(select(func.count()).select_from(Alpha))
result["favorites"] = await db.scalar(
select(func.count()).select_from(Research).where(Research.favorite.is_(True))
)
result["last_sync"] = await db.scalar(select(func.max(Alpha.synced_at)))
return result
return await Business(db).get_alpha_facets()
@api.get(
"/alphas/export",
@@ -350,106 +318,41 @@ def create_app(settings=None, wq_client=None):
)
@api.patch("/alphas/research/bulk", response_model=BulkOutput, tags=["alphas"])
async def bulk(body: BulkInput):
async with sessions() as db:
rows = (
await db.scalars(
select(Research).where(Research.alpha_id.in_(body.alpha_ids)).with_for_update()
)
).all()
if len(rows) != len(body.alpha_ids):
raise HTTPException(404, "部分 Alpha 尚未同步,本次未修改任何记录")
for row in rows:
tags = (set(row.tags) | set(body.add_tags)) - set(body.remove_tags)
try:
await set_tags(db, row, list(tags))
except ValueError as exc:
raise HTTPException(422, str(exc)) from None
if body.state:
row.state = body.state
row.updated_at = now()
await db.commit()
return {"updated": len(rows)}
async def bulk(body: BulkUpdate):
async with sessions.begin() as db:
return await Business(db).bulk_update_research(body)
@api.get("/alphas/{alpha_id}", response_model=AlphaDetail, tags=["alphas"])
async def detail(alpha_id: str):
async with sessions() as db:
a, r = await db.get(Alpha, alpha_id), await db.get(Research, alpha_id)
if a is None:
raise HTTPException(404, "Alpha 尚未同步")
return {
**summary(a, r),
**{
key: getattr(a, key)
for key in (
"expression",
"selection",
"combo",
"settings",
"is_metrics",
"os_metrics",
"checks",
)
},
}
return await Business(db).get_alpha(alpha_id)
@api.patch("/alphas/{alpha_id}/research", response_model=OkOutput, tags=["alphas"])
async def research(alpha_id: str, body: ResearchInput):
async with sessions() as db:
row = await db.scalar(select(Research).where(Research.alpha_id == alpha_id).with_for_update())
if row is None:
raise HTTPException(404, "Alpha 尚未同步")
for key, value in body.model_dump(exclude_unset=True).items():
if key == "tags":
await set_tags(db, row, value)
else:
setattr(row, key, value)
row.updated_at = now()
await db.commit()
return {"ok": True}
async def research(alpha_id: str, body: ResearchUpdate):
async with sessions.begin() as db:
return await Business(db).update_research(alpha_id, body)
@api.get("/alphas/{alpha_id}/pnl", response_model=PnlOutput, tags=["alphas"])
async def pnl(alpha_id: str):
async with sessions() as db:
if not await db.get(Alpha, alpha_id):
raise HTTPException(404, "Alpha 尚未同步")
record = await db.get(Pnl, alpha_id)
return {
"cached": record is not None,
"points": record.points if record else [],
"fetched_at": record.fetched_at if record else None,
}
return await Business(db).get_alpha_pnl(alpha_id)
@api.post("/sync-jobs", status_code=202, response_model=JobOutput, tags=["sync-jobs"])
async def new_job(body: JobInput):
async with sessions() as db:
# Serialize creation against the singleton account, avoiding duplicate full scans.
account = await db.scalar(select(Account).where(Account.id == 1).with_for_update())
if not account.password_encrypted or account.connection_status in ("disconnected", "error"):
raise HTTPException(409, "请先连接 WorldQuant")
payload = {"alpha_ids": body.alpha_ids}
existing = (
await db.scalars(select(Job).where(Job.kind == body.kind, Job.status.in_(ACTIVE)))
).all()
for job in existing:
if job.payload == payload:
return job
job = await create_job(db, body.kind, payload)
runner.wake.set()
return job
async with sessions.begin() as db:
result = await Business(db).create_sync_job(body)
await notify_job(runner, "create_sync_job", result)
return result
@api.get("/sync-jobs", response_model=list[JobOutput], tags=["sync-jobs"])
async def get_jobs():
async with sessions() as db:
return (await db.scalars(select(Job).order_by(Job.created_at.desc()).limit(100))).all()
return await Business(db).list_jobs()
@api.get("/sync-jobs/{job_id}", response_model=JobOutput, tags=["sync-jobs"])
async def get_job(job_id: str):
async with sessions() as db:
job = await db.get(Job, job_id)
if job is None:
raise HTTPException(404, "任务不存在")
return job
return await Business(db).get_job_status(job_id)
@api.get("/sync-jobs/{job_id}/errors", response_model=list[JobErrorOutput], tags=["sync-jobs"])
async def job_errors(job_id: str):
@@ -461,38 +364,18 @@ def create_app(settings=None, wq_client=None):
@api.post("/sync-jobs/{job_id}/cancel", response_model=OkOutput, tags=["sync-jobs"])
async def cancel_job(job_id: str):
async with sessions() as db:
job = await db.get(Job, job_id)
if job is None:
raise HTTPException(404, "任务不存在")
if job.status not in ACTIVE:
return {"ok": True}
job.cancel_requested = True
if job.status != "running":
job.status = "cancelled"
await db.commit()
await runner.cancel(job_id)
return {"ok": True}
async with sessions.begin() as db:
result = await Business(db).cancel_job(job_id)
await notify_job(runner, "cancel_job", result)
return result
@api.post("/sync-jobs/{job_id}/retry", response_model=JobOutput, tags=["sync-jobs"])
async def retry_job(job_id: str):
async with sessions() as db:
job = await db.get(Job, job_id)
if job is None:
raise HTTPException(404, "任务不存在")
if job.status not in (
"failed",
"cancelled",
"completed_with_errors",
"waiting_connection",
"waiting_auth",
):
raise HTTPException(409, "该任务当前不需要重试")
job.status, job.error, job.cancel_requested, job.next_retry_at = "queued", None, False, None
job.updated_at = now()
await db.commit()
runner.wake.set()
return job
async with sessions.begin() as db:
result = await Business(db).retry_job(job_id)
await notify_job(runner, "retry_job", result)
return result
app.include_router(api)
app.include_router(ai_router(ai_runtime))
return app