```python
#!/usr/bin/env python3
"""Anomaly Detection Monitor for Adversarial ML Attacks.

This script provides a practical, vendor-neutral starting point for monitoring
requests to an ML model for suspicious patterns that may indicate adversarial
activity.

What it does:
- Loads request telemetry from CSV or JSONL
- Builds a baseline from known-normal records
- Scores each request using a simple, explainable anomaly model
- Combines feature, prediction, and context signals
- Emits a CSV report and optional JSON summary

What it does not do:
- It does not block traffic by default
- It does not claim to detect every attack
- It does not replace secure model training, drift monitoring, or human review

Expected input fields are flexible, but the following columns are useful:
- request_id: unique request identifier
- timestamp: ISO-8601 timestamp or similar
- source_id: user, client, device, or API identity
- source_ip: source IP address
- feature_1..feature_n: numeric feature columns
- prediction_confidence: model confidence score between 0 and 1
- prediction_entropy: optional entropy-like score
- top_margin: difference between top 1 and top 2 class scores
- request_count_5m: request burst metric
- duplicate_score: similarity to recent requests

Example usage:
  python anomaly_monitor.py --input requests.csv --output anomalies.csv --baseline baseline.csv
  python anomaly_monitor.py --input requests.jsonl --output anomalies.csv --threshold 0.75
"""

from __future__ import annotations

import argparse
import csv
import json
import math
import statistics
import sys
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple


DEFAULT_NUMERIC_EXCLUDE = {
    "prediction_confidence",
    "prediction_entropy",
    "top_margin",
    "request_count_5m",
    "duplicate_score",
}


@dataclass
class ScoreResult:
    request_id: str
    anomaly_score: float
    label: str
    reasons: List[str]


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Score model requests for adversarial-ML-related anomalies using simple baseline statistics."
    )
    parser.add_argument("--input", required=True, help="Path to CSV or JSONL file containing request telemetry.")
    parser.add_argument(
        "--baseline",
        help="Optional path to CSV or JSONL file with normal traffic used to build the baseline. If omitted, the input file is used.",
    )
    parser.add_argument("--output", required=True, help="Path to write the scored CSV report.")
    parser.add_argument(
        "--summary",
        help="Optional path to write a JSON summary containing counts and threshold information.",
    )
    parser.add_argument(
        "--threshold",
        type=float,
        default=0.75,
        help="Anomaly score threshold for labeling a request as suspicious (default: 0.75).",
    )
    parser.add_argument(
        "--min-baseline-records",
        type=int,
        default=25,
        help="Minimum number of records required to compute a usable baseline (default: 25).",
    )
    return parser.parse_args()


def read_records(path: str) -> List[Dict[str, Any]]:
    file_path = Path(path)
    if not file_path.exists():
        raise FileNotFoundError(f"Input file not found: {path}")
    if file_path.suffix.lower() == ".csv":
        return read_csv(file_path)
    if file_path.suffix.lower() in {".jsonl", ".json"}:
        return read_jsonl(file_path)
    raise ValueError(f"Unsupported file type: {file_path.suffix}. Use .csv, .jsonl, or .json")


def read_csv(path: Path) -> List[Dict[str, Any]]:
    with path.open("r", newline="", encoding="utf-8") as f:
        reader = csv.DictReader(f)
        return [normalize_record(row) for row in reader]


def read_jsonl(path: Path) -> List[Dict[str, Any]]:
    records: List[Dict[str, Any]] = []
    with path.open("r", encoding="utf-8") as f:
        text = f.read().strip()
        if not text:
            return records
        if text.startswith("["):
            loaded = json.loads(text)
            if not isinstance(loaded, list):
                raise ValueError("JSON input must be a list of objects.")
            for item in loaded:
                if isinstance(item, dict):
                    records.append(normalize_record(item))
            return records
        f.seek(0)
        for line in f:
            line = line.strip()
            if not line:
                continue
            obj = json.loads(line)
            if isinstance(obj, dict):
                records.append(normalize_record(obj))
    return records


def normalize_record(record: Dict[str, Any]) -> Dict[str, Any]:
    normalized: Dict[str, Any] = {}
    for k, v in record.items():
        key = str(k).strip()
        if isinstance(v, str):
            value = v.strip()
            if value == "":
                normalized[key] = None
                continue
            normalized[key] = coerce_value(value)
        else:
            normalized[key] = v
    return normalized


def coerce_value(value: str) -> Any:
    try:
        if value.lower() in {"true", "false"}:
            return value.lower() == "true"
        if "." in value or "e" in value.lower():
            return float(value)
        return int(value)
    except Exception:
        return value


def safe_float(value: Any) -> Optional[float]:
    try:
        if value is None:
            return None
        if isinstance(value, bool):
            return None
        return float(value)
    except Exception:
        return None


def extract_numeric_columns(records: Sequence[Dict[str, Any]]) -> List[str]:
    candidates: Dict[str, int] = defaultdict(int)
    for record in records:
        for key, value in record.items():
            if key in {"request_id", "timestamp", "source_id", "source_ip", "label"}:
                continue
            if safe_float(value) is not None:
                candidates[key] += 1
    numeric_cols = [k for k, count in candidates.items() if count > 0]
    return numeric_cols


@dataclass
class BaselineStats:
    means: Dict[str, float]
    stdevs: Dict[str, float]
    medians: Dict[str, float]
    mads: Dict[str, float]
    feature_columns: List[str]


def build_baseline(records: Sequence[Dict[str, Any]], min_records: int) -> BaselineStats:
    if len(records) < min_records:
        raise ValueError(
            f"Not enough baseline records: {len(records)} provided, need at least {min_records}."
        )

    numeric_cols = extract_numeric_columns(records)
    feature_cols = [c for c in numeric_cols if c not in DEFAULT_NUMERIC_EXCLUDE]

    means: Dict[str, float] = {}
    stdevs: Dict[str, float] = {}
    medians: Dict[str, float] = {}
    mads: Dict[str, float] = {}

    for col in feature_cols:
        values = [safe_float(r.get(col)) for r in records]
        values = [v for v in values if v is not None]
        if not values:
            continue
        means[col] = statistics.fmean(values)
        stdevs[col] = statistics.pstdev(values) if len(values) > 1 else 1.0
        medians[col] = statistics.median(values)
        abs_dev = [abs(v - medians[col]) for v in values]
        mads[col] = statistics.median(abs_dev) if abs_dev else 1.0

    return BaselineStats(means=means, stdevs=stdevs, medians=medians, mads=mads, feature_columns=feature_cols)


def zscore(value: float, mean: float, stdev: float) -> float:
    if stdev <= 1e-12:
        return 0.0
    return abs(value - mean) / stdev


def robust_score(value: float, median: float, mad: float) -> float:
    scale = 1.4826 * mad if mad > 1e-12 else 1.0
    return abs(value - median) / scale


def clamp01(x: float) -> float:
    return max(0.0, min(1.0, x))


def evaluate_record(record: Dict[str, Any], baseline: BaselineStats) -> ScoreResult:
    reasons: List[str] = []
    scores: List[float] = []

    # Feature-space anomaly score
    feature_scores: List[float] = []
    for col in baseline.feature_columns:
        val = safe_float(record.get(col))
        if val is None:
            continue
        mean = baseline.means.get(col)
        stdev = baseline.stdevs.get(col)
        median = baseline.medians.get(col)
        mad = baseline.mads.get(col)
        if mean is None or stdev is None or median is None or mad is None:
            continue
        feature_scores.append(clamp01(max(zscore(val, mean, stdev), robust_score(val, median, mad)) / 6.0))

    if feature_scores:
        feature_anomaly = sum(feature_scores) / len(feature_scores)
        scores.append(feature_anomaly)
        if feature_anomaly >= 0.5:
            reasons.append("feature-space deviation from baseline")

    # Prediction-layer signals
    confidence = safe_float(record.get("prediction_confidence"))
    entropy = safe_float(record.get("prediction_entropy"))
    top_margin = safe_float(record.get("top_margin"))

    if confidence is not None:
        low_conf = clamp01(1.0 - confidence)
        scores.append(low_conf)
        if confidence < 0.4:
            reasons.append("low model confidence")

    if entropy is not None:
        scores.append(clamp01(entropy / 3.0))
        if entropy > 1.5:
            reasons.append("high output uncertainty")

    if top_margin is not None:
        scores.append(clamp01(1.0 - top_margin))
        if top_margin < 0.2:
            reasons.append("small margin between top classes")

    # Behavioral / contextual signals
    request_count = safe_float(record.get("request_count_5m"))
    duplicate_score = safe_float(record.get("duplicate_score"))

    if request_count is not None:
        scores.append(clamp01(request_count / 20.0))
        if request_count >= 10:
            reasons.append("burst of requests in short window")

    if duplicate_score is not None:
        scores.append(clamp01(duplicate_score))
        if duplicate_score >= 0.7:
            reasons.append("near-duplicate request pattern")

    # Simple contextual heuristics
    source_id = str(record.get("source_id") or "").strip()
    source_ip = str(record.get("source_ip") or "").strip()
    if source_id and source_ip and source_id.count("|") > 2:
        scores.append(0.2)
    if source_ip and source_ip.lower() in {"unknown", "0.0.0.0"}:
        scores.append(0.25)
        reasons.append("unusual or missing source identity")

    if not scores:
        final_score = 0.0
    else:
        final_score = clamp01(sum(scores) / len(scores))

    label = "suspicious" if final_score >= 0.75 else "normal"
    request_id = str(record.get("request_id") or record.get("id") or "")
    if not request_id:
        request_id = "<missing-request-id>"

    return ScoreResult(request_id=request_id, anomaly_score=round(final_score, 4), label=label, reasons=sorted(set(reasons)))


def write_report(path: str, results: Sequence[ScoreResult]) -> None:
    with Path(path).open("w", newline="", encoding="utf-8") as f:
        writer = csv.writer(f)
        writer.writerow(["request_id", "anomaly_score", "label", "reasons"])
        for r in results:
            writer.writerow([r.request_id, f"{r.anomaly_score:.4f}", r.label, "; ".join(r.reasons)])


def write_summary(path: str, results: Sequence[ScoreResult], threshold: float, source_count: int) -> None:
    suspicious = sum(1 for r in results if r.label == "suspicious")
    summary = {
        "input_records": source_count,
        "scored_records": len(results),
        "suspicious_records": suspicious,
        "normal_records": len(results) - suspicious,
        "threshold": threshold,
        "generated_at_utc": datetime.utcnow().isoformat(timespec="seconds") + "Z",
    }
    Path(path).write_text(json.dumps(summary, indent=2), encoding="utf-8")


def main() -> int:
    args = parse_args()
    try:
        input_records = read_records(args.input)
        baseline_records = read_records(args.baseline) if args.baseline else input_records
        baseline = build_baseline(baseline_records, args.min_baseline_records)

        results = [evaluate_record(record, baseline) for record in input_records]
        # Respect the user-specified threshold by re-labeling results when needed.
        adjusted: List[ScoreResult] = []
        for r in results:
            label = "suspicious" if r.anomaly_score >= args.threshold else "normal"
            adjusted.append(ScoreResult(r.request_id, r.anomaly_score, label, r.reasons))

        write_report(args.output, adjusted)
        if args.summary:
            write_summary(args.summary, adjusted, args.threshold, len(input_records))

        suspicious = sum(1 for r in adjusted if r.label == "suspicious")
        print(f"Processed {len(input_records)} records. Suspicious: {suspicious}. Output: {args.output}")
        return 0
    except Exception as exc:
        print(f"Error: {exc}", file=sys.stderr)
        return 1


if __name__ == "__main__":
    raise SystemExit(main())
```