strix/strix/report/state.py

1227 lines
48 KiB
Python

import json
import logging
import re
import subprocess
import threading
from collections.abc import Callable
from datetime import UTC, datetime
from importlib.metadata import PackageNotFoundError, version
from pathlib import Path
from typing import TYPE_CHECKING, Any, Optional, cast
from uuid import uuid4
from strix.config import codex
from strix.config.loader import load_settings
from strix.core.paths import run_dir_for, runtime_state_dir
from strix.report.coverage import write_coverage
from strix.report.pricing import resolve_litellm_model
from strix.report.sarif import write_sarif
from strix.report.writer import (
read_run_record,
write_executive_report,
write_run_record,
write_vulnerabilities,
)
from strix.telemetry import posthog, scarf
if TYPE_CHECKING:
from agents.usage import Usage
logger = logging.getLogger(__name__)
_global_report_state: Optional["ReportState"] = None
_CONTROL_CHARS = re.compile(r"[\x00-\x1f\x7f]+")
class ReportRollbackError(RuntimeError):
"""Raised when a failed deletion could not rewrite the on-disk indexes.
The report is still on file in memory; the indexes are rewritten from
memory on the next save.
"""
def __init__(self, report_id: str, *, cause: BaseException) -> None:
self.report_id = report_id
self.cause = cause
super().__init__(
f"Deletion of report '{report_id}' failed ({cause}) and the on-disk indexes "
"could not be restored; the report is still on file and the indexes are "
"rewritten on the next save"
)
def _strix_version() -> str | None:
"""Best-effort package version for the SARIF tool.driver.version field."""
try:
return version("strix-agent")
except PackageNotFoundError:
return None
# Content a revision may replace. The identity of the finding (id, timestamp,
# finding_class) and its original author stay put. dependency_metadata is
# replaced whole, so a caller carries the package identity over itself.
UPDATABLE_REPORT_FIELDS = frozenset(
{
"title",
"dependency_metadata",
"severity",
"description",
"impact",
"target",
"technical_analysis",
"poc_description",
"poc_script_code",
"remediation_steps",
"evidence",
"assumptions",
"counterevidence",
"confidence",
"validation_status",
"confidence_rationale",
"severity_change_conditions",
"fix_effort",
"cvss",
"cvss_breakdown",
"endpoint",
"method",
"cve",
"cwe",
"code_locations",
"http_exchange_ids",
"fix_verification",
"fix_pr_body",
"fix_candidate",
"fix_preparation",
}
)
_LOWERCASE_REPORT_FIELDS = frozenset({"severity", "confidence", "fix_effort"})
# Fields that only describe another field. A revision may raise the rating or
# replace the locations without restating the reasoning behind the old one, and
# that leftover reasoning then contradicts the finding it annotates
# ("confidence: high" beside a rationale calling the evidence unconfirmed). When
# the field they describe changes and the update carries no replacement, they
# are dropped rather than kept.
_DEPENDENT_REPORT_FIELDS: dict[str, tuple[str, ...]] = {
"confidence": ("confidence_rationale",),
"severity": ("severity_change_conditions",),
"cvss": ("cvss_breakdown",),
"code_locations": ("fix_verification",),
}
def _clean_title(title: str) -> str:
"""Return a single-line finding title.
A title quotes text from the scanned target, so it can carry newlines, tabs or
other control characters. Those break every artifact that renders the title on
one line, such as the markdown heading, the CSV cell and the TUI list. Control
characters become spaces and runs of whitespace collapse to one space.
"""
return " ".join(_CONTROL_CHARS.sub(" ", title).split())
def _number(value: Any) -> int | float:
try:
return float(value or 0)
except (TypeError, ValueError):
return 0
def _parse_repo_full_name(uri: str) -> str | None:
"""Extract ``owner/repo`` from a git URL or slug, else None."""
text = uri.strip().removesuffix(".git")
if not text:
return None
if "@" in text and ":" in text.split("@", 1)[1]:
# scp-style: git@host:owner/repo
text = text.split("@", 1)[1].split(":", 1)[1]
elif "://" in text:
# https://host/owner/repo
host_and_path = text.split("://", 1)[1]
text = host_and_path.split("/", 1)[1] if "/" in host_and_path else host_and_path
parts = [p for p in text.split("/") if p]
if len(parts) >= 2:
return "/".join(parts[-2:])
return None
def _git_head(repo_path: str) -> tuple[str | None, str | None]:
"""Best-effort ``(commit_sha, branch)`` for a cloned repo, or ``(None, None)``.
Used to populate SARIF versionControlProvenance. Failures (missing git,
non-repo path, detached HEAD, timeout) degrade to None so the SARIF
emit is never blocked by a provenance lookup.
"""
path = Path(repo_path)
if not path.is_dir():
return None, None
def _run(args: list[str]) -> str | None:
try:
result = subprocess.run( # noqa: S603
["git", "-C", str(path), *args], # noqa: S607
capture_output=True,
text=True,
check=False,
timeout=5,
)
except (OSError, subprocess.SubprocessError):
return None
if result.returncode != 0:
return None
return result.stdout.strip() or None
commit = _run(["rev-parse", "HEAD"])
branch = _run(["rev-parse", "--abbrev-ref", "HEAD"])
if branch == "HEAD": # detached HEAD carries no branch name
branch = None
return commit, branch
def get_global_report_state() -> Optional["ReportState"]:
return _global_report_state
def set_global_report_state(report_state: Optional["ReportState"]) -> None:
global _global_report_state # noqa: PLW0603
_global_report_state = report_state
# New run: drop any streamed-cost entries a prior run left unconsumed.
streamed_openrouter_costs.clear()
class ReportState:
"""Per-scan product artifact state plus artifact writer.
The Agents SDK owns model/tool execution, tracing, and conversation
persistence. This store keeps only Strix-owned scan artifacts and
report metadata. Live UI projections belong to the interface layer.
It does not consume SDK tracing processors.
"""
def __init__(self, run_name: str | None = None):
self.run_name = run_name
self.run_id = run_name or f"run-{uuid4().hex[:8]}"
self.start_time = datetime.now(UTC).isoformat()
self.process_start_time = self.start_time
self.end_time: str | None = None
self.vulnerability_reports: list[dict[str, Any]] = []
self.final_scan_result: str | None = None
self.scan_results: dict[str, Any] | None = None
self.scan_config: dict[str, Any] | None = None
# Imported here so importing this module never enters the agents SDK
# package (which the warm-up thread may be initializing concurrently).
from strix.report.usage import LLMUsageLedger
self._llm_usage = LLMUsageLedger()
self._telemetry_llm_usage_baseline: dict[str, Any] = {}
auth_mode = codex.auth_mode(load_settings().llm.model)
self._llm_usage.zero_cost = auth_mode == "subscription"
self.run_record: dict[str, Any] = {
"run_id": self.run_id,
"run_name": self.run_name,
"start_time": self.start_time,
"end_time": None,
"status": "running",
"auth_mode": auth_mode,
"targets_info": [],
"llm_usage": self._build_llm_usage_record(),
}
self._run_dir: Path | None = None
self._saved_vuln_ids: set[str] = set()
self.caido_url: str | None = None
self.fix_finding_callback: Callable[[dict[str, Any]], None] | None = None
self.vulnerability_found_callback: Callable[[dict[str, Any]], None] | None = None
self.vulnerability_updated_callback: Callable[[dict[str, Any]], None] | None = None
self.vulnerability_deleted_callback: Callable[[dict[str, Any]], None] | None = None
self._sarif_repo_ctx: dict[str, Any] | None = None
self._sarif_repo_ctx_ready: bool = False
self.posthog_scan_ended_sent: bool = False
self.scarf_scan_ended_sent: bool = False
self.scan_ended_exit_reason: str | None = None
def get_run_dir(self) -> Path:
if self._run_dir is None:
run_dir_name = self.run_name if self.run_name else self.run_id
self._run_dir = run_dir_for(run_dir_name)
self._run_dir.mkdir(parents=True, exist_ok=True)
return self._run_dir
def hydrate_from_run_dir(self) -> None:
"""Reload prior-scan state from ``{run_dir}/`` for resume.
Restores:
- ``vulnerability_reports`` from ``vulnerabilities.json`` so
:meth:`add_vulnerability_report` doesn't allocate a colliding
``vuln-0001`` and overwrite the prior on-disk MD.
- ``run_record`` from ``run.json`` so timestamps, run inputs,
status, and final report state have one public source of truth.
Idempotent on missing files (fresh runs land here too via the
same code path). **Raises on corruption** — silently swallowing
a corrupt ``vulnerabilities.json`` would let the next vuln
allocate ``vuln-0001`` and overwrite the prior MD on disk
(data loss). Caller is expected to fail the run loud and let
the user inspect ``{run_dir}`` or pick a fresh ``--run-name``.
"""
run_dir = self.get_run_dir()
data = read_run_record(run_dir)
if data:
self.run_record.update(data)
if isinstance(data.get("start_time"), str):
self.start_time = data["start_time"]
if isinstance(data.get("end_time"), str):
self.end_time = data["end_time"]
scan_results = data.get("scan_results")
if isinstance(scan_results, dict):
self.scan_results = scan_results
self.final_scan_result = self._format_final_scan_result(scan_results)
self._hydrate_llm_usage(data.get("llm_usage"))
self._telemetry_llm_usage_baseline = self._build_llm_usage_record()
logger.info("report state hydrated run.json from %s", run_dir)
json_path = run_dir / "vulnerabilities.json"
if json_path.exists():
try:
data = json.loads(json_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise RuntimeError(
f"vulnerabilities.json at {json_path} is corrupt ({exc}); "
f"refusing to start fresh — that would overwrite prior "
f"vulnerability MDs on disk. Inspect or delete the run dir.",
) from exc
if not isinstance(data, list):
raise RuntimeError(
f"vulnerabilities.json at {json_path} is not a list",
)
self.vulnerability_reports = [r for r in data if isinstance(r, dict)]
for r in self.vulnerability_reports:
# A finding written before the class was persisted still carries the
# metadata of its class, so name the class it always had.
if not r.get("finding_class"):
r["finding_class"] = (
"dependency_cve" if r.get("dependency_metadata") else "dynamic"
)
title = r.get("title")
stale_md = False
if isinstance(title, str):
r["title"] = _clean_title(title)
stale_md = r["title"] != title
rid = r.get("id")
# A finding already on disk keeps its markdown, unless cleaning
# changed the title: the heading on disk then needs a rewrite.
if isinstance(rid, str) and not stale_md:
self._saved_vuln_ids.add(rid)
logger.info(
"report state hydrated %d vulnerability report(s)",
len(self.vulnerability_reports),
)
def add_vulnerability_report(
self,
title: str,
severity: str,
description: str | None = None,
impact: str | None = None,
target: str | None = None,
technical_analysis: str | None = None,
poc_description: str | None = None,
poc_script_code: str | None = None,
remediation_steps: str | None = None,
evidence: str | None = None,
assumptions: str | None = None,
counterevidence: str | None = None,
confidence: str | None = None,
confidence_rationale: str | None = None,
severity_change_conditions: str | None = None,
fix_effort: str | None = None,
cvss: float | None = None,
cvss_breakdown: dict[str, str] | None = None,
endpoint: str | None = None,
method: str | None = None,
cve: str | None = None,
cwe: str | None = None,
code_locations: list[dict[str, Any]] | None = None,
http_exchange_ids: list[str] | None = None,
fix_verification: str | None = None,
fix_pr_body: str | None = None,
validation_status: str = "unconfirmed",
finding_class: str | None = None,
dependency_metadata: dict[str, str] | None = None,
fix_candidate: dict[str, Any] | None = None,
fix_preparation: dict[str, Any] | None = None,
agent_id: str | None = None,
agent_name: str | None = None,
) -> str:
report_id = self._next_report_id()
report: dict[str, Any] = {
"id": report_id,
"title": _clean_title(title),
"severity": severity.lower().strip(),
"timestamp": datetime.now(UTC).strftime("%Y-%m-%d %H:%M:%S UTC"),
}
if description:
report["description"] = description.strip()
if impact:
report["impact"] = impact.strip()
if target:
report["target"] = target.strip()
if technical_analysis:
report["technical_analysis"] = technical_analysis.strip()
if poc_description:
report["poc_description"] = poc_description.strip()
if poc_script_code:
report["poc_script_code"] = poc_script_code.strip()
if remediation_steps:
report["remediation_steps"] = remediation_steps.strip()
if evidence:
report["evidence"] = evidence.strip()
if assumptions:
report["assumptions"] = assumptions.strip()
if counterevidence:
report["counterevidence"] = counterevidence.strip()
if confidence:
report["confidence"] = confidence.strip().lower()
if confidence_rationale:
report["confidence_rationale"] = confidence_rationale.strip()
if severity_change_conditions:
report["severity_change_conditions"] = severity_change_conditions.strip()
if fix_effort:
report["fix_effort"] = fix_effort.strip().lower()
if cvss is not None:
report["cvss"] = cvss
if cvss_breakdown:
report["cvss_breakdown"] = cvss_breakdown
if endpoint:
report["endpoint"] = endpoint.strip()
if method:
report["method"] = method.strip()
if cve:
report["cve"] = cve.strip()
if cwe:
report["cwe"] = cwe.strip()
if code_locations:
report["code_locations"] = code_locations
if http_exchange_ids:
report["http_exchange_ids"] = http_exchange_ids
if fix_verification:
report["fix_verification"] = fix_verification.strip()
if fix_pr_body:
report["fix_pr_body"] = fix_pr_body.strip()
report["validation_status"] = validation_status
report["finding_class"] = (finding_class or "dynamic").strip().lower()
if dependency_metadata:
report["dependency_metadata"] = dependency_metadata
if fix_candidate:
report["fix_candidate"] = fix_candidate
if fix_preparation:
report["fix_preparation"] = fix_preparation
if agent_id:
report["agent_id"] = agent_id
if agent_name:
report["agent_name"] = agent_name
if self.vulnerability_found_callback:
self.vulnerability_found_callback(report)
self.vulnerability_reports.append(report)
logger.info(f"Added vulnerability report: {report_id} - {title}")
posthog.finding(severity, cwe=cwe, is_cve=bool(cve))
scarf.finding(severity, cwe=cwe, is_cve=bool(cve))
self.save_run_data()
self._notify_fix(report)
return report_id
def _deleted_vulnerability_reports(self) -> list[dict[str, Any]]:
raw = self.run_record.get("deleted_vulnerability_reports")
return [e for e in raw if isinstance(e, dict)] if isinstance(raw, list) else []
def _next_report_id(self) -> str:
"""Allocate the id after every id this run has ever handed out.
A deleted report leaves the list, so counting entries would hand its id
to the next finding and let that finding overwrite the deleted MD on disk
and inherit its history in every consumer that keys on the id.
"""
used = 0
for entry in [*self.vulnerability_reports, *self._deleted_vulnerability_reports()]:
match = re.fullmatch(r"vuln-(\d+)", str(entry.get("id", "")))
if match:
used = max(used, int(match.group(1)))
return f"vuln-{used + 1:04d}"
def update_vulnerability_report(
self,
report_id: str,
fields: dict[str, Any],
*,
update_reason: str | None = None,
updated_by_agent_id: str | None = None,
updated_by_agent_name: str | None = None,
) -> dict[str, Any] | None:
"""Apply a revision to an existing report, keeping its id.
A field that only describes a field this update replaces is dropped when
the update carries no replacement for it, so the revised report cannot
state a new rating beside the superseded reasoning for the old one.
Returns the revised report, or ``None`` when the id is unknown or when
nothing in ``fields`` changes it.
"""
report = next((r for r in self.vulnerability_reports if r.get("id") == report_id), None)
if report is None:
logger.warning("cannot update unknown vulnerability report %s", report_id)
return None
changed: dict[str, Any] = {}
for key, raw_value in fields.items():
if key not in UPDATABLE_REPORT_FIELDS or (raw_value is None and key != "fix_candidate"):
continue
value = raw_value
if isinstance(value, str):
value = _clean_title(value) if key == "title" else value.strip()
if key in _LOWERCASE_REPORT_FIELDS:
value = value.lower()
if not value:
continue
if report.get(key) == value:
continue
changed[key] = value
superseded = {
dependent
for primary, dependents in _DEPENDENT_REPORT_FIELDS.items()
if primary in changed
for dependent in dependents
if dependent not in changed and report.get(dependent) not in (None, "", [], {})
}
if not changed and not superseded:
logger.info("update for %s carried no new content; keeping it as is", report_id)
return None
entry: dict[str, Any] = {
"timestamp": datetime.now(UTC).strftime("%Y-%m-%d %H:%M:%S UTC"),
"fields": sorted(changed),
}
if superseded:
entry["dropped_fields"] = sorted(superseded)
if update_reason and update_reason.strip():
entry["reason"] = update_reason.strip()[:500]
if updated_by_agent_id:
entry["agent_id"] = updated_by_agent_id
if updated_by_agent_name:
entry["agent_name"] = updated_by_agent_name
for key in ("severity", "cvss", "confidence"):
if key in changed and report.get(key) is not None:
entry[f"previous_{key}"] = report[key]
raw_history = report.get("update_history")
history: list[dict[str, Any]] = (
[e for e in raw_history if isinstance(e, dict)] if isinstance(raw_history, list) else []
)
history.append(entry)
revised = {**report, **changed}
for dependent in superseded:
revised.pop(dependent, None)
revised["update_history"] = history
revised["updated_at"] = entry["timestamp"]
# Persistence must accept the revision before local state changes. A
# failed callback leaves the old evidence intact and the update retryable.
if self.vulnerability_updated_callback:
self.vulnerability_updated_callback(revised)
report.clear()
report.update(revised)
# The markdown on disk still shows the superseded evidence, so let the
# writer re-render it.
self._saved_vuln_ids.discard(report_id)
logger.info(
"Updated vulnerability report %s (%s)",
report_id,
", ".join(entry["fields"]) or "no field replaced",
)
self.save_run_data()
self._notify_fix(report)
return report
def delete_vulnerability_report(
self,
report_id: str,
*,
delete_reason: str,
deleted_by_agent_id: str | None = None,
deleted_by_agent_name: str | None = None,
) -> dict[str, Any] | None:
"""Withdraw a report from the run, keeping a record of the withdrawal.
Returns the removed report, or ``None`` when the id is unknown. Every
agent in a run shares one trust boundary, so any of them may withdraw
any report (just as any of them may revise one); the run record keeps
who filed it, who withdrew it and why. The report leaves
``vulnerability_reports`` and its rendered artifacts, and its id is
never reissued.
The rewritten artifacts and the ``vulnerability_deleted_callback`` must
both accept the deletion first: if either fails the report is put back
and the error propagates, so the deletion can be retried
(:class:`ReportRollbackError` when the indexes could not be put back).
"""
report = next((r for r in self.vulnerability_reports if r.get("id") == report_id), None)
if report is None:
logger.warning("cannot delete unknown vulnerability report %s", report_id)
return None
entry: dict[str, Any] = {
"id": report_id,
"title": report.get("title"),
"severity": report.get("severity"),
"filed_at": report.get("timestamp"),
"deleted_at": datetime.now(UTC).strftime("%Y-%m-%d %H:%M:%S UTC"),
"reason": delete_reason.strip()[:500],
}
if report.get("agent_id"):
entry["filed_by_agent_id"] = report["agent_id"]
if deleted_by_agent_id:
entry["agent_id"] = deleted_by_agent_id
if deleted_by_agent_name:
entry["agent_name"] = deleted_by_agent_name
position = self.vulnerability_reports.index(report)
was_saved = report_id in self._saved_vuln_ids
history = self._deleted_vulnerability_reports()
self.vulnerability_reports.remove(report)
self._saved_vuln_ids.discard(report_id)
self.run_record["deleted_vulnerability_reports"] = [*history, entry]
try:
# Local artifacts first: they can be put back if persistence then
# refuses, whereas a row deleted elsewhere cannot.
self._sync_llm_usage_record()
self._write_artifacts()
if self.vulnerability_deleted_callback:
self.vulnerability_deleted_callback({**report, "deletion": entry})
except Exception as exc:
self.vulnerability_reports.insert(position, report)
if was_saved:
self._saved_vuln_ids.add(report_id)
if history:
self.run_record["deleted_vulnerability_reports"] = history
else:
self.run_record.pop("deleted_vulnerability_reports", None)
try:
self._write_artifacts()
except Exception as rollback_exc:
logger.exception(
"could not restore artifacts after failed deletion of %s", report_id
)
raise ReportRollbackError(report_id, cause=exc) from rollback_exc
raise
md_path = self.get_run_dir() / "vulnerabilities" / f"{report_id}.md"
try:
md_path.unlink(missing_ok=True)
except OSError:
logger.exception("could not remove %s", md_path)
self._notify_fix({**report, "deletion": entry})
logger.info("Deleted vulnerability report %s - %s", report_id, report.get("title"))
return report
def _notify_fix(self, report: dict[str, Any]) -> None:
if self.fix_finding_callback:
try:
self.fix_finding_callback(dict(report))
except Exception:
logger.exception("Could not schedule fix for %s", report.get("id"))
def get_existing_vulnerabilities(self) -> list[dict[str, Any]]:
return list(self.vulnerability_reports)
def record_sdk_usage(
self,
*,
agent_id: str,
usage: "Usage | None",
agent_name: str | None = None,
model: str | None = None,
) -> None:
"""Record SDK-native token usage for one completed model run/cycle."""
if self._llm_usage.record(
agent_id=agent_id,
agent_name=agent_name,
model=model,
usage=usage,
):
self.save_run_data()
def record_observed_llm_cost(self, cost: float) -> None:
self._llm_usage.record_observed_cost(cost)
def record_llm_provider(
self,
provider: str,
*,
agent_id: str | None,
input_tokens: int,
cached_tokens: int,
cost: float,
) -> None:
self._llm_usage.record_provider(
provider,
agent_id=agent_id,
input_tokens=input_tokens,
cached_tokens=cached_tokens,
cost=cost,
cache_block_tokens=load_settings().llm.cache_block_tokens,
)
def get_process_llm_providers(self) -> dict[str, dict[str, float]]:
"""Per-provider usage since this process started, like get_process_llm_usage."""
baseline = self._telemetry_llm_usage_baseline.get("providers") or {}
providers: dict[str, dict[str, float]] = {}
for name, tally in (self._llm_usage.to_record().get("providers") or {}).items():
before = baseline.get(name) or {}
delta = {key: max(0, value - _number(before.get(key))) for key, value in tally.items()}
if delta["requests"]:
providers[name] = delta
return providers
def get_total_llm_usage(self) -> dict[str, Any]:
return dict(self.run_record.get("llm_usage") or self._build_llm_usage_record())
def get_process_llm_usage(self) -> dict[str, int | float]:
"""Return LLM usage accumulated since this process started."""
usage = self._llm_usage.to_record()
return {
key: max(
0, _number(usage.get(key)) - _number(self._telemetry_llm_usage_baseline.get(key))
)
for key in ("requests", "input_tokens", "output_tokens", "total_tokens", "cost")
}
def get_process_duration_seconds(self) -> float:
"""Return this process's elapsed wall time for telemetry."""
try:
start = datetime.fromisoformat(self.process_start_time.replace("Z", "+00:00"))
duration = (datetime.now(start.tzinfo) - start).total_seconds()
return max(0.0, duration)
except (ValueError, TypeError, AttributeError):
return 0.0
def get_total_llm_cost(self) -> float:
"""Live accumulated LLM cost, independent of the persisted run-record snapshot."""
return self._llm_usage.total_cost
def update_scan_final_fields(
self,
executive_summary: str,
methodology: str,
technical_analysis: str,
recommendations: str,
) -> None:
self.scan_results = {
"scan_completed": True,
"executive_summary": executive_summary.strip(),
"methodology": methodology.strip(),
"technical_analysis": technical_analysis.strip(),
"recommendations": recommendations.strip(),
"success": True,
}
self.final_scan_result = self._format_final_scan_result(self.scan_results)
self.run_record["scan_results"] = self.scan_results
logger.info("Updated scan final fields")
self.run_record["assessment_completed_at"] = datetime.now(UTC).isoformat()
self.save_run_data(mark_complete=self.fix_finding_callback is None)
if self.fix_finding_callback is None:
posthog.end(self, exit_reason="finished_by_tool")
scarf.end(self, exit_reason="finished_by_tool")
def record_mcp_connections(self, names: list[str]) -> None:
"""Note the MCP servers this run connected, and persist it.
Saved as soon as the run connects rather than at the end, so an interface
reading the record mid-run can already attribute a tool call to the
server it went out to.
"""
if self.run_record.get("mcp_connections") == names:
return
self.run_record["mcp_connections"] = names
self.save_run_data()
def record_mcp_connection_status(self, status: list[dict[str, Any]]) -> None:
"""Persist the run's non-secret MCP connection status roster.
``status`` is one entry per connection carrying only ``name``,
``provider``, ``tool_count``, and ``dead`` (no config, url, token, or
auth). Saved as soon as the run connects and rewritten each time a
connection dies, so the viewer, which rebuilds its display by re-reading
the run's files from disk, can show a live connections panel and health
without any in-memory event sink. Kept separate from the
``mcp_connections`` name list so neither field repurposes the other.
"""
if self.run_record.get("mcp_connection_status") == status:
return
self.run_record["mcp_connection_status"] = status
self.save_run_data()
def set_scan_config(self, config: dict[str, Any]) -> None:
self.scan_config = config
self.run_record["status"] = "running"
self.run_record["end_time"] = None
self.run_record.pop("scan_results", None)
self.end_time = None
self.scan_results = None
self.final_scan_result = None
self.run_record.update(
{
"targets_info": config.get("targets", []),
"instruction": config.get("user_instructions", ""),
"scan_mode": config.get("scan_mode", "deep"),
"diff_scope": config.get("diff_scope", {"active": False}),
"non_interactive": bool(config.get("non_interactive", False)),
"local_sources": config.get("local_sources", []),
"scope_mode": config.get("scope_mode", "auto"),
"diff_base": config.get("diff_base"),
}
)
def save_run_data(self, mark_complete: bool = False, status: str | None = None) -> None:
if mark_complete:
self.end_time = datetime.now(UTC).isoformat()
self.run_record["end_time"] = self.end_time
self.run_record["status"] = "completed"
elif status and self.run_record.get("status") != "completed":
current_status = self.run_record.get("status")
if status == "stopped" and current_status in {"failed", "interrupted"}:
status = str(current_status)
if self.end_time is None:
self.end_time = datetime.now(UTC).isoformat()
self.run_record["end_time"] = self.end_time
self.run_record["status"] = status
self._sync_llm_usage_record()
self._save_artifacts()
def cleanup(self, status: str = "stopped") -> None:
self.save_run_data(status=status)
def _format_final_scan_result(self, scan_results: dict[str, Any]) -> str:
return f"""# Executive Summary
{str(scan_results.get("executive_summary", "")).strip()}
# Methodology
{str(scan_results.get("methodology", "")).strip()}
# Technical Analysis
{str(scan_results.get("technical_analysis", "")).strip()}
# Recommendations
{str(scan_results.get("recommendations", "")).strip()}
"""
def _coverage_document(self) -> dict[str, Any] | None:
"""Assemble the coverage record, or None when it can't be built.
Coverage is a secondary artifact: a failure here must not cost the
caller its findings, so this swallows and logs rather than raising
into :meth:`_save_artifacts`.
"""
try:
from strix.report.coverage import build_coverage_document, read_agent_graph
from strix.tools.coverage.tools import get_coverage_entries
return build_coverage_document(
run_record=self.run_record,
entries=get_coverage_entries(),
agent_graph=read_agent_graph(runtime_state_dir(self.get_run_dir())),
vulnerability_reports=self.vulnerability_reports,
exit_reason=self.scan_ended_exit_reason,
)
except Exception:
logger.exception("coverage document build failed (non-fatal)")
return None
def _save_artifacts(self) -> None:
"""Write scan artifacts under ``run_dir``; a write failure is logged."""
try:
self._write_artifacts()
except (OSError, RuntimeError):
logger.exception("Failed to save scan data")
def _write_artifacts(self) -> None:
"""Write scan artifacts under ``run_dir``, raising when the index or
run record cannot be written."""
run_dir = self.get_run_dir()
run_dir.mkdir(parents=True, exist_ok=True)
coverage = self._coverage_document()
if coverage is not None:
try:
write_coverage(run_dir, coverage)
except OSError:
logger.exception("coverage.json write failed (non-fatal)")
if self.final_scan_result:
write_executive_report(run_dir, self.final_scan_result)
# An index is written for an empty list too once a report was
# deleted, or the CSV/JSON on disk would still list it.
if self.vulnerability_reports or self._deleted_vulnerability_reports():
write_vulnerabilities(run_dir, self.vulnerability_reports, self._saved_vuln_ids)
# SARIF 2.1.0 emitter for CI / ASPM integration. Always emit (even
# empty) so a clean run overwrites a prior findings.sarif rather than
# leaving a stale one — codeql-action's "absent from new submission →
# fixed" needs the fresh empty doc to auto-resolve alerts. Isolated
# in its own try: a SARIF-build error must NEVER break the CSV/MD/
# run-record path (the emitter's own contract).
try:
write_sarif(
run_dir,
self.vulnerability_reports,
tool_version=_strix_version(),
repository_context=self._sarif_repository_context(),
coverage=coverage,
)
except Exception:
logger.exception("SARIF emit failed (non-fatal; CSV/MD unaffected)")
write_run_record(run_dir, self.run_record)
logger.info("Essential scan data saved to: %s", run_dir)
def _sarif_repository_context(self) -> dict[str, Any] | None:
"""Repo/commit/branch context for SARIF provenance (repo scans only).
Cached after first derivation — ``_save_artifacts`` runs on every
state save, and the git lookup only needs to happen once per run.
Returns None for URL / IP (DAST) targets that have no repository.
"""
if not self._sarif_repo_ctx_ready:
self._sarif_repo_ctx = self._derive_repository_context()
self._sarif_repo_ctx_ready = True
return self._sarif_repo_ctx
def _derive_repository_context(self) -> dict[str, Any] | None:
targets = self.run_record.get("targets_info") or []
if not isinstance(targets, list):
return None
repo_targets = [
target
for target in targets
if isinstance(target, dict) and target.get("type") in {"repository", "local_code"}
]
# Provenance binds the whole run to one repo; with multiple repo targets
# that's ambiguous, so omit it rather than mis-attributing later repos'
# findings to the first repo's URI/commit.
if len(repo_targets) != 1:
return None
target = repo_targets[0]
details = target.get("details") or {}
if not isinstance(details, dict):
return None
uri = details.get("target_repo")
if target.get("type") == "local_code" and details.get("target_path"):
uri = Path(details["target_path"]).resolve().as_uri()
if not isinstance(uri, str) or not uri.strip():
return None
context: dict[str, Any] = {"repositoryUri": uri.strip()}
full_name = _parse_repo_full_name(uri)
if full_name:
context["repositoryFullName"] = full_name
cloned = details.get("cloned_repo_path") or details.get("target_path")
if isinstance(cloned, str) and cloned.strip():
commit, branch = _git_head(cloned.strip())
if commit:
context["commitSha"] = commit
if branch:
context["branch"] = branch
context["ref"] = f"refs/heads/{branch}"
return context
def get_repository_context(self) -> dict[str, Any] | None:
return self._derive_repository_context()
def _sync_llm_usage_record(self) -> None:
self.run_record["llm_usage"] = self._build_llm_usage_record()
def _build_llm_usage_record(self) -> dict[str, Any]:
return self._llm_usage.to_record()
def _hydrate_llm_usage(self, raw_usage: Any) -> None:
self._llm_usage.hydrate(raw_usage)
self._sync_llm_usage_record()
def openrouter_stream_cost(usage: Any) -> float | None:
"""Total OpenRouter-reported cost from a raw stream ``usage`` block, or None.
Non-BYOK responses bill everything to ``usage.cost``. BYOK responses put the
OpenRouter fee in ``usage.cost`` (often 0) and the provider charge in
``usage.cost_details.upstream_inference_cost``, so BYOK totals sum the two.
"""
if not isinstance(usage, dict):
return None
total = 0.0
cost = usage.get("cost")
if isinstance(cost, int | float) and cost > 0:
total += float(cost)
if bool(usage.get("is_byok")):
details = usage.get("cost_details")
upstream = details.get("upstream_inference_cost") if isinstance(details, dict) else None
if isinstance(upstream, int | float) and upstream > 0:
total += float(upstream)
return total if total > 0 else None
def _response_id(completion_response: Any) -> str | None:
response_id = getattr(completion_response, "id", None)
if response_id is None and isinstance(completion_response, dict):
response_id = cast("dict[str, Any]", completion_response).get("id")
return response_id if isinstance(response_id, str) and response_id else None
class StreamedOpenRouterCosts:
"""Correlates OpenRouter's per-stream cost from the parser to the cost callback.
LiteLLM rebuilds streamed responses from token-only chunks and drops the
``usage.cost`` OpenRouter reports in its final stream chunk (its non-streamed
path preserves it; streaming snapshots hidden params at stream start). Every
scan streams, so the OpenRouter streaming handler (see strix.config.models)
records the cost here keyed by response id, and the callback takes it back out
for the matching rebuilt response. Entries are removed on read; ``clear()``
runs per scan so nothing accumulates across runs.
"""
def __init__(self) -> None:
self._costs: dict[str, float] = {}
self._lock = threading.Lock()
def remember(self, response_id: Any, usage: Any) -> None:
cost = openrouter_stream_cost(usage)
if cost is None or not (isinstance(response_id, str) and response_id):
return
with self._lock:
self._costs[response_id] = cost
def take(self, completion_response: Any) -> float | None:
response_id = _response_id(completion_response)
if response_id is None:
return None
with self._lock:
return self._costs.pop(response_id, None)
def clear(self) -> None:
with self._lock:
self._costs.clear()
streamed_openrouter_costs = StreamedOpenRouterCosts()
def record_openrouter_provider(provider: Any, usage: Any) -> None:
"""Tally which upstream provider served a stream, from its final usage chunk.
OpenRouter spreads one model across many providers whose prices, quantization
and prompt caching differ, so this is what shows where a scan's tokens went.
"""
# Deferred: request_log pulls in the agents SDK, which strix.report must not import.
from strix.llm.request_log import current_call_context
report_state = get_global_report_state()
if report_state is None or not isinstance(usage, dict):
return
details = usage.get("prompt_tokens_details")
report_state.record_llm_provider(
provider if isinstance(provider, str) and provider else "unknown",
agent_id=current_call_context().agent_id,
input_tokens=int(_number(usage.get("prompt_tokens"))),
cached_tokens=int(_number(details.get("cached_tokens")))
if isinstance(details, dict)
else 0,
cost=openrouter_stream_cost(usage) or 0.0,
)
def litellm_cost_callback(
kwargs: Any,
completion_response: Any,
_start_time: Any = None,
_end_time: Any = None,
) -> None:
"""LiteLLM ``success_callback`` adapter; forwards observed cost to the active scan."""
cost: float | None = None
raw = kwargs.get("response_cost") if isinstance(kwargs, dict) else None
if isinstance(raw, int | float) and raw > 0:
cost = float(raw)
if cost is None:
hidden = getattr(completion_response, "_hidden_params", None) or {}
candidate = hidden.get("response_cost") if isinstance(hidden, dict) else None
if isinstance(candidate, int | float) and candidate > 0:
cost = float(candidate)
else:
headers = hidden.get("additional_headers") or {} if isinstance(hidden, dict) else {}
raw = (
headers.get("llm_provider-x-litellm-response-cost")
if isinstance(headers, dict)
else None
)
try:
value = float(raw) if raw is not None else None
except (TypeError, ValueError):
value = None
if value is not None and value > 0:
cost = value
if cost is None:
cost = _usage_reported_cost(completion_response)
# Recover the exact OpenRouter cost the streaming handler stashed for this
# response — LiteLLM drops it from streamed usage, so nothing above sees it.
if cost is None:
cost = streamed_openrouter_costs.take(completion_response)
if cost is None:
cost = _estimate_response_cost(kwargs, completion_response)
if cost is None or cost <= 0:
return
report_state = get_global_report_state()
if report_state is None:
return
try:
report_state.record_observed_llm_cost(cost)
except Exception:
logger.exception("Failed to record observed LiteLLM cost")
def _usage_reported_cost(completion_response: Any) -> float | None:
"""Provider-reported cost from the ``usage`` block (e.g. OpenRouter).
Non-BYOK responses charge everything to ``usage.cost``. BYOK responses
charge only the OpenRouter fee to ``usage.cost`` (often 0) and report the
provider charge in ``usage.cost_details.upstream_inference_cost``, so the
true BYOK total is the sum of the two.
"""
usage: Any = getattr(completion_response, "usage", None)
if usage is None and isinstance(completion_response, dict):
usage = cast("dict[str, Any]", completion_response).get("usage")
if usage is None:
return None
def _field(container: Any, name: str) -> Any:
if isinstance(container, dict):
return cast("dict[str, Any]", container).get(name)
return getattr(container, name, None)
total = 0.0
usage_cost = _field(usage, "cost")
if isinstance(usage_cost, int | float) and usage_cost > 0:
total += float(usage_cost)
if bool(_field(usage, "is_byok")):
upstream = _field(_field(usage, "cost_details"), "upstream_inference_cost")
if isinstance(upstream, int | float) and upstream > 0:
total += float(upstream)
return total if total > 0 else None
def _estimate_response_cost(kwargs: Any, completion_response: Any) -> float | None:
"""Best-effort LiteLLM cost-map estimate when no provider-reported cost exists.
LiteLLM strips provider cost fields when rebuilding streamed responses and
returns no ``response_cost`` for models missing from its cost map, so try
the provider-prefixed name, the raw name, and the bare model name.
"""
from litellm import completion_cost
model = kwargs.get("model") if isinstance(kwargs, dict) else None
if not isinstance(model, str) or not model:
if isinstance(completion_response, dict):
model = cast("dict[str, Any]", completion_response).get("model")
else:
model = getattr(completion_response, "model", None)
if not isinstance(model, str) or not model:
return None
provider = None
litellm_params = kwargs.get("litellm_params") if isinstance(kwargs, dict) else None
if isinstance(litellm_params, dict):
provider = litellm_params.get("custom_llm_provider")
usage_payload = _usage_payload(completion_response)
if usage_payload is None:
return None
candidates: list[str] = []
if isinstance(provider, str) and provider and not model.startswith(f"{provider}/"):
candidates.append(f"{provider}/{model}")
candidates.append(model)
if "/" in model:
candidates.append(model.rsplit("/", 1)[-1])
for candidate in candidates:
resolved = resolve_litellm_model(candidate)
if not resolved:
continue
try:
value = completion_cost(
completion_response={"model": resolved, "usage": usage_payload},
model=resolved,
)
except Exception: # nosec B112 # noqa: BLE001, S112
continue
if isinstance(value, int | float) and value > 0:
return float(value)
return None
def _usage_payload(completion_response: Any) -> dict[str, Any] | None:
"""Token counts as a plain dict, detached from the response's provider metadata."""
usage: Any = getattr(completion_response, "usage", None)
if usage is None and isinstance(completion_response, dict):
usage = cast("dict[str, Any]", completion_response).get("usage")
if usage is None:
return None
if hasattr(usage, "model_dump"):
usage = usage.model_dump()
if not isinstance(usage, dict):
return None
payload = cast("dict[str, Any]", usage)
if not payload.get("total_tokens") and not (
payload.get("prompt_tokens") or payload.get("completion_tokens")
):
return None
return payload