feat: add AI research chatbot with confirmed business tools
This commit is contained in:
+39
-156
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user