netx/netx_api/ne_collect_runner.py
oliver fbf73cbaaf Disable target paging after hop login and keep CLI echo on timeouts.
LLDP discover and config collect were timing out on --More--; send vendor paging-off after nested/bastion auth, and attach session-log tails so failures show where the CLI stuck.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-08-29 12:30:33 +08:00

338 lines
13 KiB
Python

from __future__ import annotations
import logging
import re
import time
import traceback
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from pathlib import Path
from typing import Any
from .collection_job_state import finalize_collection_job, sync_job_progress
from .config import settings
from .db import SessionLocal
from fastapi import HTTPException
from .cli_creds import cli_creds_skip_reason
from .cli_budget import clamp_cli_workers
from .cli_resolve import resolve_cli_target
from .models import NeCollectionJob, NeCollectionRun
from .ne_collection_paths import clear_run_output_files, run_output_dir
from .ne_crypto import CredentialCryptoError
from .ne_netmiko import send_show_command
from .ne_session_factory import close_netmiko_connection, open_netmiko_connection
_log = logging.getLogger("netx.ne.collect")
_executor: ThreadPoolExecutor | None = None
def _executor_pool() -> ThreadPoolExecutor:
global _executor
if _executor is None:
workers = clamp_cli_workers(int(settings.ne_collect_max_workers or 8))
_executor = ThreadPoolExecutor(max_workers=workers, thread_name_prefix="ne-collect")
return _executor
def shutdown_ne_collect_executor(*, wait: bool = False) -> None:
global _executor
if _executor is not None:
try:
_executor.shutdown(wait=wait, cancel_futures=True)
except TypeError:
_executor.shutdown(wait=wait)
_executor = None
def _format_run_error(exc: BaseException) -> str:
head = f"{type(exc).__name__}: {exc}"
# RuntimeError from format_cli_failure already carries session log — keep it.
if isinstance(exc, RuntimeError) and "--- session log ---" in str(exc):
return str(exc)[:4000]
tb = traceback.format_exc().strip()
text = f"{head}\n{tb}" if tb else head
return text[:4000]
def _safe_filename_part(text: str) -> str:
s = re.sub(r'[<>:"/\\|?*]', "_", str(text or "").strip())
return s[:80] or "device"
def _collect_on_device(
creds: dict[str, Any],
commands: list[str],
*,
read_timeout_sec: int | None = None,
conn_holder: dict[str, Any] | None = None,
) -> str:
import io
from .ne_cli_errors import format_cli_failure, session_log_text
from .ne_netmiko import disable_target_paging
per_cmd = int(read_timeout_sec if read_timeout_sec is not None else (settings.ne_collect_read_timeout_sec or 120))
session_timeout = per_cmd * max(1, len(commands)) + 60
log_buf = io.BytesIO()
chunks: list[str] = []
conn = None
try:
try:
conn = open_netmiko_connection(
creds, session_timeout=session_timeout, session_log=log_buf
)
except Exception as exc:
raise RuntimeError(format_cli_failure(exc, session_log_text(log_buf), limit=4000)) from exc
if conn_holder is not None:
conn_holder["conn"] = conn
conn_holder["session_log"] = log_buf
# Belt-and-suspenders: hop/bastion nested CLIs and missed session_prep.
try:
disable_target_paging(
conn,
vendor=str(creds.get("vendor") or ""),
device_type=str(creds.get("device_type") or ""),
)
except Exception:
_log.debug("collection paging disable failed", exc_info=True)
prompt = str(conn.find_prompt() or "")
for command in commands:
if conn_holder is not None and conn_holder.get("timed_out"):
raise TimeoutError("collection_aborted")
ts = datetime.now().isoformat(timespec="seconds")
chunks.append(f'>>> [{ts}] {{"String":"{command}", "Match":"{prompt}", "Timeout":0}}\n')
try:
out = send_show_command(conn, command, read_timeout=per_cmd)
except Exception as exc:
partial = "".join(chunks)
transcript = session_log_text(log_buf) or partial
raise RuntimeError(
format_cli_failure(exc, transcript, limit=4000)
) from exc
chunks.append(str(out or ""))
chunks.append("\n")
return "".join(chunks)
except Exception as exc:
# Surface echo for mid-command Netmiko failures not already wrapped.
if isinstance(exc, RuntimeError) and "--- session log ---" in str(exc):
raise
partial = "".join(chunks)
transcript = session_log_text(log_buf) or partial
if transcript:
raise RuntimeError(format_cli_failure(exc, transcript, limit=4000)) from exc
raise
finally:
if conn_holder is not None:
conn_holder.pop("conn", None)
close_netmiko_connection(conn)
def _collect_with_timeout(
creds: dict[str, Any],
commands: list[str],
*,
read_timeout_sec: int | None = None,
) -> str:
from .cli_timeout import run_cli_with_timeout
from .ne_cli_errors import format_cli_failure, session_log_text
per_cmd = int(read_timeout_sec if read_timeout_sec is not None else (settings.ne_collect_read_timeout_sec or 120))
cap = int(settings.ne_collect_run_timeout_cap_sec or 600)
budget = min(cap, per_cmd * max(1, len(commands)) + 90)
holder: dict[str, Any] = {}
try:
return run_cli_with_timeout(
lambda: _collect_on_device(
creds, commands, read_timeout_sec=per_cmd, conn_holder=holder
),
timeout_sec=budget,
conn_holder=holder,
label="collection",
acquire_budget=True,
)
except TimeoutError as exc:
transcript = session_log_text(holder.get("session_log"))
raise RuntimeError(format_cli_failure(exc, transcript, limit=4000)) from exc
def _update_run(run_id: str, **fields: Any) -> None:
db = SessionLocal()
try:
row = db.get(NeCollectionRun, run_id)
if not row:
return
for key, val in fields.items():
setattr(row, key, val)
db.commit()
finally:
db.close()
def _job_is_paused(job_id: str) -> bool:
db = SessionLocal()
try:
job = db.get(NeCollectionJob, job_id)
return not job or str(job.status or "") == "paused"
finally:
db.close()
def _claim_run(job_id: str, run_id: str) -> bool:
"""Atomically move a pending run to running; retry while DB rows become visible."""
for attempt in range(10):
db = SessionLocal()
try:
run = db.get(NeCollectionRun, run_id)
job = db.get(NeCollectionJob, job_id)
if not run or not job:
time.sleep(0.05 * (attempt + 1))
continue
job_status = str(job.status or "")
run_status = str(run.status or "")
if job_status == "paused":
if run_status == "pending":
run.status = "cancelled"
run.message = "paused"
run.ended_at = datetime.now()
db.commit()
return False
if job_status != "running":
return False
if run_status == "running":
return False
if run_status != "pending":
return False
run.status = "running"
run.message = "collecting"
run.started_at = datetime.now()
db.commit()
return True
finally:
db.close()
time.sleep(0.05 * (attempt + 1))
return False
def _run_single(job_id: str, run_id: str, commands: list[str]) -> None:
if not _claim_run(job_id, run_id):
db = SessionLocal()
try:
run = db.get(NeCollectionRun, run_id)
st = str(run.status or "") if run else ""
if st in ("cancelled", "success", "fail", "running"):
sync_job_progress(db, job_id)
finalize_collection_job(db, job_id)
return
if st == "pending":
_log.warning("collection claim failed job=%s run=%s", job_id, run_id)
_update_run(
run_id,
status="fail",
message="collection_claim_failed",
ended_at=datetime.now(),
)
sync_job_progress(db, job_id)
finalize_collection_job(db, job_id)
finally:
db.close()
return
db = SessionLocal()
try:
run = db.get(NeCollectionRun, run_id)
if not run:
return
if _job_is_paused(job_id):
_update_run(run_id, status="cancelled", message="paused", ended_at=datetime.now())
return
try:
source = str(getattr(run, "ne_source", None) or "managed").strip().lower()
tid = str(run.ne_id or "").strip()
if source == "ume":
creds, _device = resolve_cli_target(db, ume_ne_id=tid)
else:
creds, _device = resolve_cli_target(db, managed_ne_id=tid)
skip = cli_creds_skip_reason(creds, interactive=False)
if skip:
_update_run(run_id, status="fail", message=skip[:1020], ended_at=datetime.now())
return
output = _collect_with_timeout(creds, commands)
finished_at = datetime.now()
name_part = _safe_filename_part(str(run.ne_name or creds.get("name") or "ne"))
ip_part = _safe_filename_part(str(run.ne_ip or creds.get("ip_address") or "ip"))
ts = finished_at.strftime("%Y%m%d%H%M%S")
rel_dir = Path(job_id) / run_id
out_dir = run_output_dir(job_id, run_id)
clear_run_output_files(job_id, run_id)
out_dir.mkdir(parents=True, exist_ok=True)
filename = f"{name_part}-{ip_part}-{ts}.txt"
full_path = out_dir / filename
max_bytes = int(getattr(settings, "ne_collect_max_output_bytes", 0) or 0)
if max_bytes > 0 and len(output.encode("utf-8", errors="replace")) > max_bytes:
# Truncate by characters roughly under byte budget.
truncated = output.encode("utf-8", errors="replace")[:max_bytes]
text = truncated.decode("utf-8", errors="replace")
text += f"\n...[truncated {max_bytes} bytes cap]\n"
full_path.write_text(text, encoding="utf-8", errors="replace")
else:
full_path.write_text(output, encoding="utf-8", errors="replace")
rel_path = str(rel_dir / filename).replace("\\", "/")
_update_run(
run_id,
status="success",
message="collected",
output_rel_path=rel_path,
ended_at=finished_at,
)
except CredentialCryptoError as exc:
_update_run(run_id, status="fail", message=str(exc)[:1020], ended_at=datetime.now())
except HTTPException as exc:
detail = str(exc.detail if exc.detail is not None else exc)[:1020]
_update_run(run_id, status="fail", message=detail, ended_at=datetime.now())
except Exception as exc:
_log.exception("collection failed run=%s", run_id)
_update_run(run_id, status="fail", message=_format_run_error(exc), ended_at=datetime.now())
finally:
db.close()
db2 = SessionLocal()
try:
sync_job_progress(db2, job_id)
finalize_collection_job(db2, job_id)
finally:
db2.close()
def _run_collect_safe(job_id: str, run_id: str, commands: list[str]) -> None:
try:
_run_single(job_id, run_id, commands)
except Exception:
_log.exception("collection task crashed job=%s run=%s", job_id, run_id)
_update_run(
run_id,
status="fail",
message="collection_worker_crashed",
ended_at=datetime.now(),
)
db = SessionLocal()
try:
sync_job_progress(db, job_id)
finalize_collection_job(db, job_id)
finally:
db.close()
def schedule_collection_runs(job_id: str, run_ids: list[str], commands: list[str]) -> int:
pool = _executor_pool()
cmd_list = list(commands)
submitted = 0
for run_id in run_ids:
pool.submit(_run_collect_safe, job_id, str(run_id), cmd_list)
submitted += 1
if submitted:
_log.info("scheduled collection job=%s runs=%s", job_id, submitted)
return submitted
def dispatch_collection_runs(job_id: str, run_ids: list[str], commands: list[str]) -> int:
"""Entry point for FastAPI BackgroundTasks after the request transaction commits."""
return schedule_collection_runs(job_id, run_ids, commands)