Fix business monitoring queries, exports and collection protection

This commit is contained in:
oliver 2026-10-11 00:02:33 +08:00
parent 051a10ad95
commit 5403c7d939
4 changed files with 306 additions and 99 deletions

View file

@ -6,6 +6,7 @@ from collections import defaultdict
from datetime import datetime, timedelta
from typing import Any
from sqlalchemy import or_, select
from sqlalchemy.orm import Session
from ..models import (
@ -26,38 +27,54 @@ def _utcnow() -> datetime:
return datetime.utcnow()
def protected_batch_map(db: Session, *, task_id: str = "") -> dict[str, list[str]]:
"""batch_id → reason codes. Empty task_id = all tasks."""
def protected_batch_map(
db: Session, *, task_id: str = "", batch_ids: list[str] | None = None,
) -> dict[str, list[str]]:
"""Only read reference columns, scoped to the requested task/batches when given."""
out: dict[str, list[str]] = defaultdict(list)
wanted = set(batch_ids) if batch_ids is not None else None
if wanted == set():
return {}
scope = select(BizStateBatch.id)
q = db.query(BizStateBatch.id, BizStateBatch.is_baseline, BizStateBatch.status)
if task_id:
scope = scope.where(BizStateBatch.task_id == task_id)
q = q.filter(BizStateBatch.task_id == task_id)
if wanted is not None:
scope = scope.where(BizStateBatch.id.in_(wanted))
q = q.filter(BizStateBatch.id.in_(wanted))
batch_rows = q.all() if task_id or wanted is not None else q.filter(or_(
BizStateBatch.is_baseline.is_(True), BizStateBatch.status.in_(("queued", "running")),
)).all()
if task_id:
wanted = {row.id for row in batch_rows}
def add(bid: str, reason: str) -> None:
b = str(bid or "").strip()
if not b:
if not b or (wanted is not None and b not in wanted):
return
if reason not in out[b]:
out[b].append(reason)
q = db.query(BizStateBatch)
if task_id:
q = q.filter(BizStateBatch.task_id == task_id)
for b in q.filter(BizStateBatch.is_baseline.is_(True)).all():
add(b.id, "manual_baseline")
for b in batch_rows:
if b.is_baseline:
add(b.id, "manual_baseline")
if b.status in ("queued", "running"):
add(b.id, "active_collection")
for j in db.query(BizCompareJob).all():
add(j.before_batch_id, "compare_job_before")
add(j.after_batch_id, "compare_job_after")
for r in db.query(BizCompareRun).all():
add(r.before_batch_id, "compare_run_before")
add(r.after_batch_id, "compare_run_after")
for p in db.query(BizMigrationProject).all():
add(p.old_baseline_batch_id, "migration_old_baseline")
add(p.new_baseline_batch_id, "migration_new_baseline")
for r in db.query(BizMigrationRun).all():
add(r.old_batch_id, "migration_run_old")
add(r.new_batch_id, "migration_run_new")
references = (
(BizCompareJob.before_batch_id, BizCompareJob.after_batch_id, "compare_job_before", "compare_job_after"),
(BizCompareRun.before_batch_id, BizCompareRun.after_batch_id, "compare_run_before", "compare_run_after"),
(BizMigrationProject.old_baseline_batch_id, BizMigrationProject.new_baseline_batch_id, "migration_old_baseline", "migration_new_baseline"),
(BizMigrationRun.old_batch_id, BizMigrationRun.new_batch_id, "migration_run_old", "migration_run_new"),
)
for before, after, before_reason, after_reason in references:
refs = db.query(before, after)
if task_id or batch_ids is not None:
refs = refs.filter(or_(before.in_(scope), after.in_(scope)))
for before_id, after_id in refs:
add(before_id, before_reason)
add(after_id, after_reason)
return dict(out)
@ -67,7 +84,7 @@ def protected_batch_ids(db: Session, *, task_id: str = "") -> set[str]:
def batch_protect_info(db: Session, batch_id: str) -> dict[str, Any]:
reasons = protected_batch_map(db).get(batch_id, [])
reasons = protected_batch_map(db, batch_ids=[batch_id]).get(batch_id, [])
b = db.get(BizStateBatch, batch_id)
if b and bool(getattr(b, "is_baseline", False)) and "manual_baseline" not in reasons:
reasons = ["manual_baseline", *reasons]

View file

@ -9,8 +9,8 @@ from typing import Any
from uuid import uuid4
from fastapi import HTTPException
from sqlalchemy import String, cast, func, or_
from sqlalchemy.orm import Session
from sqlalchemy import String, case, cast, func, or_
from sqlalchemy.orm import Session, defer
from ..lldp_shared import resolve_vendor_key
from ..models import (
@ -354,13 +354,15 @@ def get_task(db: Session, task_id: str) -> dict[str, Any]:
.order_by(BizStateTaskItem.sort_order.asc())
.all()
)
bindings_by_item: dict[str, list[Any]] = {}
if items:
for binding in db.query(BizStateTaskItemBinding).filter(
BizStateTaskItemBinding.item_id.in_([it.id for it in items])
).all():
bindings_by_item.setdefault(binding.item_id, []).append(binding)
item_out = []
for it in items:
binds = (
db.query(BizStateTaskItemBinding)
.filter(BizStateTaskItemBinding.item_id == it.id)
.all()
)
binds = bindings_by_item.get(it.id, [])
item_out.append(
{
"id": it.id,
@ -373,6 +375,11 @@ def get_task(db: Session, task_id: str) -> dict[str, Any]:
"bindings": [{"placeholder": b.placeholder, "value": b.value} for b in binds],
}
)
return {**_task_summary(task), "items": item_out}
def _task_summary(task: BizStateTask) -> dict[str, Any]:
"""Task metadata without command items or bindings (safe for frequent polling)."""
return {
"id": task.id,
"source": task.source,
@ -396,10 +403,16 @@ def get_task(db: Session, task_id: str) -> dict[str, Any]:
if task.last_collect_ended_at
else None,
"last_error": task.last_error,
"items": item_out,
}
def get_task_progress(db: Session, task_id: str) -> dict[str, Any]:
task = db.get(BizStateTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="task_not_found")
return _task_summary(task)
def list_tasks(db: Session, *, purpose: str | None = None) -> list[dict[str, Any]]:
from sqlalchemy import or_
@ -447,6 +460,8 @@ def delete_task(db: Session, task_id: str) -> None:
task = db.get(BizStateTask, task_id)
if not task:
raise HTTPException(status_code=404, detail="task_not_found")
if task.collect_running:
raise HTTPException(status_code=409, detail="task_collecting")
batches = db.query(BizStateBatch).filter(BizStateBatch.task_id == task_id).all()
# Refuse if any batch is still referenced by compare/migration (manual baseline is OK to drop with task)
pmap = protected_batch_map(db, task_id=task_id)
@ -500,7 +515,7 @@ def list_batches(db: Session, task_id: str, *, limit: int = 50) -> list[dict[str
.limit(max(1, min(500, int(limit))))
.all()
)
pmap = protected_batch_map(db, task_id=task_id)
pmap = protected_batch_map(db, batch_ids=[b.id for b in rows])
out: list[dict[str, Any]] = []
for b in rows:
reasons = list(pmap.get(b.id, []))
@ -622,12 +637,29 @@ def get_batch(db: Session, batch_id: str) -> dict[str, Any]:
b = db.get(BizStateBatch, batch_id)
if not b:
raise HTTPException(status_code=404, detail="batch_not_found")
cmds = (
db.query(BizStateBatchCommand)
.filter(BizStateBatchCommand.batch_id == batch_id)
.order_by(BizStateBatchCommand.created_at.asc())
.all()
raw = func.coalesce(BizStateBatchCommand.raw_text, "")
raw_len = func.length(raw)
# Compute legacy counts in SQL; workbook summaries never transfer raw CLI text.
line_count = case(
(BizStateBatchCommand.raw_line_count > 0, BizStateBatchCommand.raw_line_count),
(raw_len == 0, 0),
else_=raw_len - func.length(func.replace(raw, "\n", ""))
+ case((func.substr(raw, raw_len, 1) == "\n", 0), else_=1),
)
whitespace = " \t\r\n\v\f"
trimmed = (func.btrim(raw, whitespace) if db.get_bind().dialect.name == "postgresql"
else func.trim(raw, whitespace))
cmd_raw_stats: dict[str, tuple[int, bool]] = {}
cmds = []
for c, lines, has_raw in (
db.query(BizStateBatchCommand, line_count, func.length(trimmed) > 0)
.options(defer(BizStateBatchCommand.raw_text))
.filter(BizStateBatchCommand.batch_id == batch_id)
.order_by(BizStateBatchCommand.created_at.asc(), BizStateBatchCommand.id.asc())
.all()
):
cmds.append(c)
cmd_raw_stats[c.id] = (int(lines or 0), bool(has_raw))
protect = batch_protect_info(db, batch_id)
# Per-metric row counts (generic table)
@ -668,11 +700,14 @@ def get_batch(db: Session, batch_id: str) -> dict[str, Any]:
# Prefer primary collect rows over aux / aux_cached for the same CLI
sheet_cmds[id_].append(cmd_info)
primary_cmds_by_cli: dict[str, dict[str, Any]] = {}
# Aux failures must remain visible even if their primary command was stored later.
primary_cmds_by_cli = {
normalize_command(str(c.raw_command or "")) for c in cmds
if not str(c.parse_status or "").strip().lower().startswith("aux")
}
for c in cmds:
status = str(c.parse_status or "").strip().lower()
is_aux = status.startswith("aux")
raw = c.raw_text or ""
info = {
"id": c.id,
"profile_id": c.profile_id,
@ -682,15 +717,13 @@ def get_batch(db: Session, batch_id: str) -> dict[str, Any]:
"params": c.params_json or {},
"parse_status": c.parse_status,
"row_count": c.row_count,
"raw_line_count": _cmd_raw_line_count(c),
"raw_line_count": cmd_raw_stats[c.id][0],
"declared_total": _cmd_declared_total(c),
"message": c.message,
"has_raw": bool(str(raw).strip()),
"has_raw": cmd_raw_stats[c.id][1],
"is_aux": is_aux,
}
cmd_n = normalize_command(str(c.raw_command or ""))
if not is_aux and cmd_n:
primary_cmds_by_cli.setdefault(cmd_n, info)
# Commands sheet: hide successful aux when the same CLI already has a primary row;
# keep failed/skipped aux visible so partial reasons are not hidden.
if is_aux and cmd_n and cmd_n in primary_cmds_by_cli:
@ -746,10 +779,10 @@ def get_batch(db: Session, batch_id: str) -> dict[str, Any]:
"params": c.params_json or {},
"parse_status": c.parse_status,
"row_count": c.row_count,
"raw_line_count": _cmd_raw_line_count(c),
"raw_line_count": cmd_raw_stats[c.id][0],
"declared_total": _cmd_declared_total(c),
"message": c.message,
"has_raw": bool(str(c.raw_text or "").strip()),
"has_raw": cmd_raw_stats[c.id][1],
"is_aux": True,
}
)
@ -844,6 +877,7 @@ def list_batch_metric_rows(
size_n = max(1, min(200, int(page_size or 50)))
kw_n = str(kw or "").strip()
col_n = str(column or "").strip()
like = "%" + kw_n.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%"
fields = metric_field_map().get(mid) or []
columns = [
@ -859,30 +893,30 @@ def list_batch_metric_rows(
if mid == "lldp_neighbor":
q = db.query(BizStateLldpNeighbor).filter(BizStateLldpNeighbor.batch_id == batch_id)
if kw_n:
like = f"%{kw_n}%"
if col_n == "local_if":
q = q.filter(BizStateLldpNeighbor.local_if.ilike(like))
q = q.filter(BizStateLldpNeighbor.local_if.ilike(like, escape="\\"))
elif col_n == "remote_sys":
q = q.filter(BizStateLldpNeighbor.remote_sys.ilike(like))
q = q.filter(BizStateLldpNeighbor.remote_sys.ilike(like, escape="\\"))
elif col_n == "remote_if":
q = q.filter(BizStateLldpNeighbor.remote_if.ilike(like))
q = q.filter(BizStateLldpNeighbor.remote_if.ilike(like, escape="\\"))
elif col_n == "remote_ip":
q = q.filter(BizStateLldpNeighbor.remote_ip.ilike(like))
q = q.filter(BizStateLldpNeighbor.remote_ip.ilike(like, escape="\\"))
elif col_n == "protocol":
q = q.filter(BizStateLldpNeighbor.protocol.ilike(like))
q = q.filter(BizStateLldpNeighbor.protocol.ilike(like, escape="\\"))
else:
q = q.filter(
or_(
BizStateLldpNeighbor.local_if.ilike(like),
BizStateLldpNeighbor.remote_sys.ilike(like),
BizStateLldpNeighbor.remote_if.ilike(like),
BizStateLldpNeighbor.remote_ip.ilike(like),
BizStateLldpNeighbor.protocol.ilike(like),
BizStateLldpNeighbor.local_if.ilike(like, escape="\\"),
BizStateLldpNeighbor.remote_sys.ilike(like, escape="\\"),
BizStateLldpNeighbor.remote_if.ilike(like, escape="\\"),
BizStateLldpNeighbor.remote_ip.ilike(like, escape="\\"),
BizStateLldpNeighbor.protocol.ilike(like, escape="\\"),
)
)
total = int(q.count() or 0)
rows_db = (
q.order_by(BizStateLldpNeighbor.local_if.asc())
q.order_by(BizStateLldpNeighbor.local_if.asc(), BizStateLldpNeighbor.remote_sys.asc(),
BizStateLldpNeighbor.remote_if.asc(), BizStateLldpNeighbor.id.asc())
.offset((page_n - 1) * size_n)
.limit(size_n)
.all()
@ -911,12 +945,11 @@ def list_batch_metric_rows(
BizStateMetricRow.metric_id == mid,
)
if kw_n:
like = f"%{kw_n}%"
if col_n:
# JSON path as text — works on Postgres JSONB and SQLite JSON
q = q.filter(cast(BizStateMetricRow.data_json[col_n], String).ilike(like))
q = q.filter(cast(BizStateMetricRow.data_json[col_n].as_string(), String).ilike(like, escape="\\"))
else:
q = q.filter(cast(BizStateMetricRow.data_json, String).ilike(like))
q = q.filter(cast(BizStateMetricRow.data_json, String).ilike(like, escape="\\"))
total = int(q.count() or 0)
rows_db = (
q.order_by(BizStateMetricRow.seq.asc(), BizStateMetricRow.id.asc())
@ -1019,56 +1052,64 @@ def export_batch_zip(db: Session, batch_id: str) -> bytes:
db.query(BizStateLldpNeighbor)
.filter(BizStateLldpNeighbor.batch_id == batch_id)
.order_by(BizStateLldpNeighbor.local_if.asc())
.all()
.yield_per(1000)
)
csv_lines = ["local_if,remote_sys,remote_if,remote_ip,protocol"]
for n in neighbors:
csv_lines.append(
",".join(
[
_csv(n.local_if),
_csv(n.remote_sys),
_csv(n.remote_if),
_csv(n.remote_ip),
_csv(n.protocol),
]
)
)
zf.writestr("tables/lldp_neighbor.csv", "\n".join(csv_lines) + "\n")
with zf.open("tables/lldp_neighbor.csv", "w") as dest:
dest.write(b"local_if,remote_sys,remote_if,remote_ip,protocol\n")
for n in neighbors:
line = ",".join(_csv(v) for v in (
n.local_if, n.remote_sys, n.remote_if, n.remote_ip, n.protocol,
)) + "\n"
dest.write(line.encode("utf-8"))
# Generic metrics CSV (stream by metric_id)
for sheet in detail.get("sheets") or []:
mid = str(sheet.get("metric_id") or "").strip()
if not mid or mid == "lldp_neighbor":
continue
rows = (
db.query(BizStateMetricRow)
.filter(
BizStateMetricRow.batch_id == batch_id,
BizStateMetricRow.metric_id == mid,
)
.order_by(BizStateMetricRow.seq.asc(), BizStateMetricRow.id.asc())
.all()
)
if not rows:
continue
recs = [dict(r.data_json or {}) for r in rows]
cols: list[str] = []
for rec in recs:
# First scan discovers the same complete, ordered header as the legacy export.
# Second scan writes rows without materializing the table or CSV.
col_keys: dict[str, None] = {}
has_rows = False
for rec in _iter_metric_export_rows(db, batch_id, mid):
has_rows = True
for k in rec.keys():
if k not in cols:
cols.append(str(k))
out_lines = [",".join(_csv(c) for c in cols)]
for rec in recs:
out_lines.append(",".join(_csv(str(rec.get(c, "") or "")) for c in cols))
col_keys.setdefault(str(k), None)
if not has_rows:
continue
cols = list(col_keys)
safe = "".join(ch if ch.isalnum() or ch in "-_" else "_" for ch in mid)[:80] or "metric"
zf.writestr(f"tables/{safe}.csv", "\n".join(out_lines) + "\n")
with zf.open(f"tables/{safe}.csv", "w") as dest:
dest.write((",".join(_csv(c) for c in cols) + "\n").encode("utf-8"))
for rec in _iter_metric_export_rows(db, batch_id, mid):
line = ",".join(_csv(rec.get(c)) for c in cols) + "\n"
dest.write(line.encode("utf-8"))
return buf.getvalue()
def _csv(v: str) -> str:
s = str(v or "")
if any(ch in s for ch in ",\"\n"):
def _iter_metric_export_rows(db: Session, batch_id: str, metric_id: str, chunk_size: int = 1000):
cursor: tuple[int, str] | None = None
while True:
query = db.query(BizStateMetricRow.seq, BizStateMetricRow.id, BizStateMetricRow.data_json).filter(
BizStateMetricRow.batch_id == batch_id, BizStateMetricRow.metric_id == metric_id,
)
if cursor is not None:
seq, row_id = cursor
query = query.filter(or_(
BizStateMetricRow.seq > seq,
(BizStateMetricRow.seq == seq) & (BizStateMetricRow.id > row_id),
))
chunk = query.order_by(BizStateMetricRow.seq.asc(), BizStateMetricRow.id.asc()).limit(chunk_size).all()
if not chunk:
return
for row in chunk:
yield row.data_json or {}
cursor = (chunk[-1].seq, chunk[-1].id)
def _csv(v: Any) -> str:
s = "" if v is None else str(v)
if any(ch in s for ch in ",\"\r\n"):
return '"' + s.replace('"', '""') + '"'
return s

View file

@ -173,6 +173,11 @@ def api_get_task(task_id: str, db: Session = Depends(get_db)) -> dict[str, Any]:
return svc.get_task(db, task_id)
@router.get("/tasks/{task_id}/progress")
def api_task_progress(task_id: str, db: Session = Depends(get_db)) -> dict[str, Any]:
return svc.get_task_progress(db, task_id)
@router.get("/tasks/{task_id}/commands")
def api_plan_task_commands(
task_id: str,