Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion api/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@

_INSECURE_JWT_DEFAULT = "change-me-in-production"
_MIN_JWT_SECRET_LENGTH = 32
_MAX_AUTHORIZATION_HEADER_LENGTH = 8192
_GENERATE_CMD = 'python -c "import secrets; print(secrets.token_urlsafe(32))"'


Expand Down Expand Up @@ -167,7 +168,7 @@ def verify_jwt() -> None:
return None

auth = request.headers.get("Authorization", "")
if not auth.startswith("Bearer "):
if len(auth) > _MAX_AUTHORIZATION_HEADER_LENGTH or not auth.startswith("Bearer "):
return jsonify(
{
"error": "Missing or malformed Authorization header",
Expand Down
5 changes: 4 additions & 1 deletion api/observability.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@

import logging
import os
import re
import time
import uuid

Expand All @@ -35,6 +36,7 @@
logger = logging.getLogger(__name__)

REQUEST_ID_HEADER = "X-Request-ID"
_REQUEST_ID_RE = re.compile(r"^[A-Za-z0-9._:-]{1,128}$")

# --------------------------------------------------------------------------- #
# Prometheus metrics #
Expand Down Expand Up @@ -160,7 +162,8 @@ def init_app(app: Flask) -> None:

@app.before_request
def _start_observability() -> None:
g.request_id = request.headers.get(REQUEST_ID_HEADER) or str(uuid.uuid4())
supplied_request_id = request.headers.get(REQUEST_ID_HEADER, "")
g.request_id = supplied_request_id if _REQUEST_ID_RE.fullmatch(supplied_request_id) else str(uuid.uuid4())
g.request_start_time = time.perf_counter()

@app.after_request
Expand Down
101 changes: 57 additions & 44 deletions api/routes/ai.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,19 @@
from api.rate_limit import rate_limit
from api.services.ai_provider import PROVIDERS as SUPPORTED_PROVIDERS
from api.services.ai_provider import get_completion
from api.validation import (
MAX_API_KEY_LENGTH,
MAX_MODEL_LENGTH,
MAX_QUESTION_LENGTH,
MODEL_RE,
VALIDATION_ERROR_MESSAGE,
ValidationError,
bounded_string,
choice,
findings_list,
reject_unknown_fields,
require_json_object,
)
from ai.retriever import retrieve, VectorStoreNotBuilt

ai_bp = Blueprint("ai", __name__)
Expand Down Expand Up @@ -148,14 +161,18 @@ def _context_for(query):


def _read_request():
body = request.get_json(silent=True)
if not body:
return None, (jsonify({"error": "Request body must be JSON"}), 400)
if not body.get("provider"):
return None, (jsonify({"error": "provider is required"}), 400)
if not body.get("api_key"):
return None, (jsonify({"error": "api_key is required"}), 400)
return body, None
try:
body = require_json_object(request.get_json(silent=True))
reject_unknown_fields(body, {"provider", "api_key", "model", "findings", "question"})
body["provider"] = choice(body.get("provider"), "provider", SUPPORTED_PROVIDERS, case="lower")
body["api_key"] = bounded_string(body.get("api_key"), "api_key", maximum=MAX_API_KEY_LENGTH)
if body.get("model") is not None:
body["model"] = bounded_string(body["model"], "model", maximum=MAX_MODEL_LENGTH, pattern=MODEL_RE)
if ".." in body["model"]:
raise ValidationError("model has an invalid format")
return body, None
except ValidationError:
return None, (jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400)


_AI_ERROR_MESSAGES = {
Expand All @@ -180,27 +197,21 @@ def _ai_error_response(exc: Exception, status: int, log_context: str):
@ai_bp.post("/api/ai/insights")
@rate_limit(_AI_RATE_LIMIT)
def insights():
data = request.get_json(silent=True)
if data is None:
return jsonify({"error": "Request body must be valid JSON"}), 400

provider = str(data.get("provider") or "").strip().lower()
api_key = str(data.get("api_key") or "").strip()
findings = data.get("findings")
question = str(data.get("question") or "").strip()

if not provider:
return jsonify({"error": "Missing required field: provider"}), 400
if provider not in SUPPORTED_PROVIDERS:
return jsonify({"error": f"Unsupported provider: {provider}"}), 400
if not api_key:
return jsonify({"error": "Missing required field: api_key"}), 400
if findings is None:
return jsonify({"error": "Missing required field: findings"}), 400
if not isinstance(findings, list):
return jsonify({"error": "findings must be a list"}), 400
if len(findings) == 0:
return jsonify({"error": "findings must not be empty"}), 400
data, error = _read_request()
if error:
return error
try:
provider = data["provider"]
api_key = data["api_key"]
findings = findings_list(data.get("findings"), required=True)
question = ""
if data.get("question") is not None:
if not isinstance(data["question"], str):
raise ValidationError("question must be a string")
if data["question"].strip():
question = bounded_string(data["question"], "question", maximum=MAX_QUESTION_LENGTH)
except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400

sorted_findings = sorted(findings, key=severity_rank, reverse=True)

Expand Down Expand Up @@ -234,9 +245,10 @@ def ai_summary():
body, error = _read_request()
if error:
return error
findings = body.get("findings", [])
if not isinstance(findings, list):
return jsonify({"error": "findings must be a list"}), 400
try:
findings = findings_list(body.get("findings"))
except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400

findings_text = _findings_to_text(findings)
try:
Expand Down Expand Up @@ -273,9 +285,10 @@ def ai_prioritise():
body, error = _read_request()
if error:
return error
findings = body.get("findings", [])
if not isinstance(findings, list):
return jsonify({"error": "findings must be a list"}), 400
try:
findings = findings_list(body.get("findings"))
except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400

findings_text = _findings_to_text(findings)
try:
Expand Down Expand Up @@ -319,16 +332,17 @@ def ai_ask():
body, error = _read_request()
if error:
return error
question = body.get("question", "")
if not question or not question.strip():
return jsonify({"error": "question is required"}), 400
try:
question = bounded_string(body.get("question"), "question", maximum=MAX_QUESTION_LENGTH)
findings = findings_list(body.get("findings"))
except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400

try:
context, sources = _context_for(question)
except VectorStoreNotBuilt as exc:
return _ai_error_response(exc, 503, "Vector store unavailable in ai_ask")

findings = body.get("findings", [])
findings_text = _findings_to_text(findings) if findings else "Not provided."

prompt = (
Expand Down Expand Up @@ -361,11 +375,10 @@ def ai_threat_simulation():
body, error = _read_request()
if error:
return error
findings = body.get("findings", [])
if not isinstance(findings, list):
return jsonify({"error": "findings must be a list"}), 400
if not findings:
return jsonify({"error": "findings must not be empty"}), 400
try:
findings = findings_list(body.get("findings"), required=True)
except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400

findings_text = _findings_to_text(findings)
try:
Expand Down
44 changes: 39 additions & 5 deletions api/routes/findings.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,19 +2,27 @@

import logging
import os
import re
from pathlib import Path
from flask import Blueprint, g, jsonify, request

from api.models.finding import DatabaseManager
from api.validation import (
CATEGORIES,
RULE_ID_RE,
SEVERITIES,
VALIDATION_ERROR_MESSAGE,
ValidationError,
bounded_string,
choice,
positive_integer,
uuid_string,
)

_PLAYBOOKS_DIR = (Path(__file__).parent.parent.parent / "playbooks" / "cli").resolve()

# Known rule_id shape, e.g. AZ-STOR-001. Anything else is rejected before it
# ever reaches the filesystem, closing off path traversal via a crafted or
# corrupted rule_id.
_RULE_ID_RE = re.compile(r"^[A-Z0-9]+(?:-[A-Z0-9]+)*$")

findings_bp = Blueprint("findings", __name__)
logger = logging.getLogger(__name__)

Expand All @@ -37,10 +45,30 @@ def list_findings():
scan_id - UUID of a specific scan
"""
try:
filters = {k: v for k, v in request.args.items() if k in ("severity", "category", "rule_id", "scan_id")}
allowed = {"severity", "category", "rule_id", "scan_id"}
unknown = set(request.args) - allowed
if unknown:
raise ValidationError(f"Unsupported query parameter: {sorted(unknown)[0]}")
for key in request.args:
if len(request.args.getlist(key)) != 1:
raise ValidationError(f"Query parameter {key} must be provided once")

filters = {}
if "severity" in request.args:
filters["severity"] = choice(request.args["severity"], "severity", SEVERITIES, case="upper")
if "category" in request.args:
filters["category"] = choice(request.args["category"], "category", CATEGORIES)
if "rule_id" in request.args:
filters["rule_id"] = bounded_string(
request.args["rule_id"].upper(), "rule_id", maximum=64, pattern=RULE_ID_RE
)
if "scan_id" in request.args:
filters["scan_id"] = uuid_string(request.args["scan_id"], "scan_id")
db = _get_db()
findings = db.get_findings(filters)
return jsonify({"count": len(findings), "findings": findings})
except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400
except Exception as exc:
logger.error("Failed to list findings: %s", exc)
return jsonify({"error": "Failed to retrieve findings"}), 500
Expand All @@ -50,11 +78,14 @@ def list_findings():
def get_finding(finding_id: int):
"""Return a single finding by its integer ID."""
try:
finding_id = positive_integer(finding_id, "finding_id")
db = _get_db()
finding = db.get_finding_by_id(finding_id)
if not finding:
return jsonify({"error": "Finding not found"}), 404
return jsonify(finding)
except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400
except Exception as exc:
logger.error("Failed to get finding %d: %s", finding_id, exc)
return jsonify({"error": "Database error"}), 500
Expand All @@ -68,6 +99,7 @@ def get_playbook(finding_id: int):
and combines it with the finding's remediation guidance and any CVE references.
"""
try:
finding_id = positive_integer(finding_id, "finding_id")
db = _get_db()
finding = db.get_finding_by_id(finding_id)
if not finding:
Expand All @@ -79,7 +111,7 @@ def get_playbook(finding_id: int):

cli_commands = []
script_path = None
if _RULE_ID_RE.match(rule_id or ""):
if RULE_ID_RE.match(rule_id or ""):
# Map rule_id (e.g. AZ-STOR-001) to script filename (fix_az_stor_001.sh)
script_name = "fix_" + rule_id.lower().replace("-", "_") + ".sh"
candidate = (_PLAYBOOKS_DIR / script_name).resolve()
Expand Down Expand Up @@ -124,6 +156,8 @@ def get_playbook(finding_id: int):
}
)

except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400
except Exception as exc:
logger.error("Failed to get playbook for finding %d: %s", finding_id, exc)
return jsonify({"error": "Failed to retrieve playbook"}), 500
20 changes: 19 additions & 1 deletion api/routes/scans.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,13 @@
from flask import Blueprint, g, jsonify, request

from api.models.finding import DatabaseManager
from api.validation import (
VALIDATION_ERROR_MESSAGE,
ValidationError,
reject_unknown_fields,
require_json_object,
uuid_string,
)
from scanner.cve_correlator import enrich_findings

scans_bp = Blueprint("scans", __name__)
Expand Down Expand Up @@ -39,11 +46,14 @@ def list_scans():
def get_scan_status(scan_id):
"""Return the details and status of a specific scan."""
try:
scan_id = uuid_string(scan_id, "scan_id")
db = _get_db()
scan = db.get_scan(scan_id)
if not scan:
return jsonify({"error": "Scan not found"}), 404
return jsonify(scan)
except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400
except Exception as exc:
logger.error("Failed to get scan status: %s", exc)
return jsonify({"error": "Database error"}), 500
Expand All @@ -59,11 +69,14 @@ def trigger_scan():
Returns 202 Accepted with the scan_id immediately.
"""
try:
body = request.get_json(silent=True) or {}
raw_body = request.get_json(silent=True)
body = {} if raw_body is None and not request.data else require_json_object(raw_body)
reject_unknown_fields(body, {"subscription_id"})
subscription_id = body.get("subscription_id") or os.environ.get("AZURE_SUBSCRIPTION_ID")

if not subscription_id:
return jsonify({"error": "subscription_id is required"}), 400
subscription_id = uuid_string(subscription_id, "subscription_id")

scan_id = str(uuid.uuid4())
logger.info("Async scan triggered for subscription %s (id: %s)", subscription_id, scan_id)
Expand All @@ -79,6 +92,8 @@ def trigger_scan():
{"scan_id": scan_id, "status": "pending", "message": "Scan has been queued and will start shortly."}
), 202

except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400
except Exception as exc:
logger.error("Critical error in trigger_scan route: %s", exc, exc_info=True)
return jsonify({"error": "Critical route failure"}), 500
Expand Down Expand Up @@ -153,6 +168,7 @@ def enrich_scan(scan_id):
rate-limited to one every ~7 seconds.
"""
try:
scan_id = uuid_string(scan_id, "scan_id")
db = _get_db()

# Check current status to avoid redundant NVD calls
Expand Down Expand Up @@ -189,6 +205,8 @@ def enrich_scan(scan_id):
}
), 202

except ValidationError:
return jsonify({"error": VALIDATION_ERROR_MESSAGE}), 400
except Exception as exc:
logger.error("Failed to start enrichment for scan %s: %s", scan_id, exc)
return jsonify({"error": "Internal server error"}), 500
Loading
Loading