145 lines
4.7 KiB
Python
145 lines
4.7 KiB
Python
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]
|