unis_manager/app/routers/ai.py

145 lines
4.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

from __future__ import annotations
import json
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session
from .. import ai_parser
from ..db import get_db
from ..models import ImportLog, Period, Project, Task, Update
from ..schemas import AICommitIn, AIParseIn
router = APIRouter()
def _resolve_project(db: Session, project_id: int | None, project_name: str | None) -> Project | None:
if project_id:
p = db.get(Project, project_id)
if not p:
raise HTTPException(404, "项目不存在")
return p
if project_name:
p = db.query(Project).filter(Project.name == project_name).one_or_none()
if not p:
p = Project(name=project_name)
db.add(p)
db.flush()
return p
return None
@router.post("/ai/parse")
def ai_parse(payload: AIParseIn, db: Session = Depends(get_db)):
project = _resolve_project(db, payload.project_id, payload.project_name)
period = db.get(Period, payload.period_id) if payload.period_id else None
parsed, mode, error = ai_parser.parse(
payload.raw_text,
project.name if project else "",
period.label if period else "",
db,
use_ai=payload.use_ai,
)
warning = parsed.pop("_warning", None)
return {
"mode": mode,
"error": error,
"warning": warning,
"parsed": parsed,
"project": project.to_dict() if project else None,
"period": period.to_dict() if period else None,
}
@router.post("/ai/commit")
def ai_commit(payload: AICommitIn, db: Session = Depends(get_db)):
project = _resolve_project(db, payload.project_id, payload.project_name)
if not project:
raise HTTPException(400, "缺少项目信息")
period = db.get(Period, payload.period_id)
if not period:
raise HTTPException(404, "周期不存在")
parsed = payload.parsed or {}
tasks = parsed.get("tasks") or []
created = []
if not tasks:
u = Update(
project_id=project.id,
period_id=period.id,
content=payload.raw_text,
summary=parsed.get("summary"),
progress=parsed.get("progress"),
status=parsed.get("status") or "in_progress",
risk=";".join(parsed.get("risks") or []) or None,
next_step=";".join(parsed.get("next_steps") or []) or None,
source="ai" if payload.mode == "ai" else "manual",
raw_text=payload.raw_text,
model=payload.model,
)
db.add(u)
created.append(u)
else:
risks = parsed.get("risks") or []
nexts = parsed.get("next_steps") or []
for idx, t in enumerate(tasks):
name = (t.get("name") or "执行事项").strip()[:255] or "执行事项"
task = (
db.query(Task)
.filter(Task.project_id == project.id, Task.name == name)
.one_or_none()
)
if not task:
task = Task(project_id=project.id, name=name, sort_order=idx)
db.add(task)
db.flush()
risk = (t.get("risk") or "").strip() or (risks[0] if idx == 0 and risks else None)
u = Update(
project_id=project.id,
task_id=task.id,
period_id=period.id,
content=(t.get("content") or "").strip() or None,
summary=t.get("summary") or parsed.get("summary"),
progress=t.get("progress"),
status=t.get("status") or parsed.get("status") or "in_progress",
risk=risk,
next_step=(nexts[0] if idx == 0 and nexts else None),
source="ai" if payload.mode == "ai" else "manual",
raw_text=payload.raw_text,
model=payload.model,
)
db.add(u)
created.append(u)
if parsed.get("progress") is not None:
project.progress = float(parsed["progress"]) or 0
if parsed.get("status"):
project.status = parsed["status"]
db.add(
ImportLog(
project_id=project.id,
period_id=period.id,
raw_input=payload.raw_text,
parsed_json=json.dumps(parsed, ensure_ascii=False),
mode=payload.mode,
model=payload.model,
)
)
db.commit()
return {
"ok": True,
"project": project.to_dict(),
"created": [u.to_dict() for u in created],
}
@router.get("/ai/logs")
def ai_logs(limit: int = 20, db: Session = Depends(get_db)):
rows = db.query(ImportLog).order_by(ImportLog.id.desc()).limit(limit).all()
return [r.to_dict() for r in rows]