Skip to content
Merged
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
115 changes: 112 additions & 3 deletions backend/app/schemas/endpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@

# Regex to find named bind parameters in Oracle SQL (:param_name).
_BIND_PARAM_RE = re.compile(r":([A-Za-z_]\w*)")
_ORACLE_ALT_QUOTE_PAIRS = {"[": "]", "{": "}", "(": ")", "<": ">"}

# Reject obvious string-interpolation patterns that bypass bind variables.
_UNSAFE_PATTERNS = [
Expand Down Expand Up @@ -50,10 +51,87 @@ class SnapshotConfigurationError(ValueError):
"""Raised when a snapshot endpoint cannot execute without request inputs."""


def _is_oracle_identifier_char(value: str) -> bool:
"""Match the identifier boundary used by Oracle alternative-quote prefixes."""
return (value.isascii() and value.isalnum()) or value in "_$#"


def _quoted_region_end(sql: str, start: int, quote: str) -> int:
"""Return the end of a standard SQL literal or quoted identifier."""
index = start + 1
while index < len(sql):
if sql[index] != quote:
index += 1
continue
if index + 1 < len(sql) and sql[index + 1] == quote:
index += 2
continue
return index + 1
return len(sql)


def _oracle_alt_quote_end(sql: str, start: int) -> int | None:
"""Return the end of a Q/NQ literal, or ``None`` when no prefix starts here."""
if start > 0 and _is_oracle_identifier_char(sql[start - 1]):
return None

quote_index: int
if sql[start] in "qQ":
quote_index = start + 1
elif sql[start] in "nN" and start + 1 < len(sql) and sql[start + 1] in "qQ":
quote_index = start + 2
else:
return None

delimiter_index = quote_index + 1
if (
quote_index >= len(sql)
or sql[quote_index] != "'"
or delimiter_index >= len(sql)
or sql[delimiter_index].isspace()
):
return None

closing_delimiter = _ORACLE_ALT_QUOTE_PAIRS.get(sql[delimiter_index], sql[delimiter_index])
closing_index = sql.find(closing_delimiter + "'", delimiter_index + 1)
return len(sql) if closing_index < 0 else closing_index + 2


def _mask_sql_non_code(sql: str) -> str:
"""Mask quoted and commented regions in one monotonic pass."""
segments: list[str] = []
code_start = 0
index = 0

while index < len(sql):
region_end: int | None = None
if sql.startswith("--", index):
newline_index = sql.find("\n", index + 2)
region_end = len(sql) if newline_index < 0 else newline_index
elif sql.startswith("/*", index):
comment_end = sql.find("*/", index + 2)
region_end = len(sql) if comment_end < 0 else comment_end + 2
elif sql[index] in "qQnN":
region_end = _oracle_alt_quote_end(sql, index)
if region_end is None and sql[index] in {"'", '"'}:
region_end = _quoted_region_end(sql, index, sql[index])

if region_end is None:
index += 1
continue

segments.append(sql[code_start:index])
segments.append(" " * (region_end - index))
index = region_end
code_start = region_end

segments.append(sql[code_start:])
return "".join(segments)


def extract_bind_params(sql: str) -> list[str]:
"""Return deduplicated bind parameter names from SQL text."""
# Exclude matches inside single-quoted string literals.
cleaned = re.sub(r"'[^']*'", "", sql)
"""Return deduplicated binds outside SQL literals, identifiers, and comments."""
cleaned = _mask_sql_non_code(sql)
return list(dict.fromkeys(_BIND_PARAM_RE.findall(cleaned)))


Expand Down Expand Up @@ -388,6 +466,13 @@ class SqlPreviewRequest(BaseModel):
connection_id: uuid.UUID
sql_text: str = Field(..., min_length=1)
params: dict[str, str | int | float | bool | None] = Field(default_factory=dict)
param_schema: dict[str, ParamDescriptor] = Field(
default_factory=dict,
description=(
"Typed descriptors for preview bind parameters. Bind-bearing SQL requires an exact "
"schema match so values can be coerced before Oracle execution."
),
)
max_rows: int = Field(10, ge=1, le=100)

@field_validator("sql_text")
Expand All @@ -398,6 +483,30 @@ def validate_sql(cls, v: str) -> str:
raise ValueError("; ".join(errors))
return v

@model_validator(mode="after")
def typed_schema_matches_bind_params(self) -> Self:
sql_params = set(extract_bind_params(self.sql_text))
Comment thread
badry-dev marked this conversation as resolved.
if not self.param_schema:
if sql_params:
raise ValueError(
"SQL preview bind parameters require typed schema descriptors for: "
f"{sorted(sql_params)}"
)
return self
Comment thread
badry-dev marked this conversation as resolved.

schema_params = set(self.param_schema)
undeclared = sql_params - schema_params
unused = schema_params - sql_params
if undeclared:
raise ValueError(
f"SQL references preview params not declared in schema: {sorted(undeclared)}"
)
if unused:
raise ValueError(
f"Preview schema declares params not referenced in SQL: {sorted(unused)}"
)
return self


class SqlPreviewResponse(BaseModel):
"""Response from SQL preview execution."""
Expand Down
54 changes: 51 additions & 3 deletions backend/app/services/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,7 @@
snapshot_covers_request,
unavailable_snapshot_filter_columns,
validate_snapshot_parameter_ranges,
validate_snapshot_rows_match_resolved_parameters,
)
from app.sql.executor import SqlExecutionError, execute_query
from app.sql.param_models import build_param_model
Expand Down Expand Up @@ -345,6 +346,7 @@ async def _serve_snapshot(
)

snapshot = None
integrity_failures: list[tuple[str, str]] = []
if not param_schema:
snapshot = snapshots[0]
else:
Expand All @@ -361,15 +363,58 @@ async def _serve_snapshot(
if candidate.job_run_id is None:
continue
job_run = job_runs_by_id.get(candidate.job_run_id)
if job_run is not None and snapshot_covers_request(
if job_run is None or not snapshot_covers_request(
filters=filters,
request_params=params,
resolved_params=job_run.resolved_params_json or {},
):
continue

if not isinstance(candidate.data, list):
integrity_failures.append(
(str(candidate.id), "Snapshot payload is not a row array.")
)
continue
candidate_data = candidate.data
if unavailable_snapshot_filter_columns(rows=candidate_data, filters=filters):
# Preserve the existing configuration-error response below. An endpoint edit
# can invalidate mappings even when the stored snapshot itself was valid.
snapshot = candidate
break
try:
validate_snapshot_rows_match_resolved_parameters(
rows=candidate_data,
filters=filters,
resolved_params=job_run.resolved_params_json or {},
)
except ValueError as exc:
integrity_failures.append((str(candidate.id), str(exc)))
continue
snapshot = candidate
break

if snapshot is None:
if integrity_failures:
self._log_snapshot_rejection(
request=request,
endpoint=endpoint,
path=path,
principal=principal,
started_at=started_at,
event="snapshot_integrity_failed",
response_status=status.HTTP_503_SERVICE_UNAVAILABLE,
details={
"snapshot_ids": [item[0] for item in integrity_failures],
"integrity_errors": [item[1] for item in integrity_failures],
},
)
return JSONResponse(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
content={
"code": "snapshot_integrity_failed",
"detail": "No retained snapshot passed integrity validation.",
},
)
self._log_snapshot_rejection(
request=request,
endpoint=endpoint,
Expand Down Expand Up @@ -438,18 +483,21 @@ def _log_snapshot_rejection(
principal: str | None,
started_at: float,
event: str,
response_status: int = status.HTTP_422_UNPROCESSABLE_ENTITY,
details: dict[str, object] | None = None,
) -> None:
"""Emit the required request context before a snapshot HTTP 422 response."""
"""Emit required request context before a snapshot rejection response."""
log.warning(
event,
request_id=resolve_request_id(request),
user=principal or "anonymous",
endpoint=path,
endpoint_id=str(endpoint.id),
status=status.HTTP_422_UNPROCESSABLE_ENTITY,
status=response_status,
duration_ms=round((time.perf_counter() - started_at) * 1000, 2),
method=request.method,
client_ip=request.client.host if request.client else None,
**(details or {}),
)

# ── Live mode ───────────────────────────────────────────────────────────
Expand Down
21 changes: 20 additions & 1 deletion backend/app/services/endpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
)
from app.services.schedule_bindings import ScheduleBindingError, resolve_schedule_parameters
from app.sql.executor import SqlExecutionError, execute_query
from app.sql.param_models import build_param_model

log = structlog.get_logger()

Expand Down Expand Up @@ -284,12 +285,30 @@ async def preview_sql(self, payload: SqlPreviewRequest) -> SqlPreviewResponse:
raise ValueError("Connection is not active.")

bind_params = extract_bind_params(payload.sql_text)
params: dict[str, object] = dict(payload.params)

if payload.param_schema:
try:
ParamModel = build_param_model(
{
name: descriptor.model_dump()
for name, descriptor in payload.param_schema.items()
},
enforce_required=True,
)
params = ParamModel.model_validate(payload.params).model_dump()
except ValidationError as exc:
first = exc.errors()[0]
field = ".".join(str(part) for part in first.get("loc", ())) or "?"
raise ValueError(
f"Invalid value for preview parameter '{field}': {first.get('msg')}"
) from exc

try:
columns, rows, duration_ms = await execute_query(
connection=conn,
sql=payload.sql_text,
params=dict(payload.params),
params=params,
max_rows=payload.max_rows,
)
except SqlExecutionError as exc:
Expand Down
10 changes: 10 additions & 0 deletions backend/app/services/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,10 @@
from app.repositories.snapshot import SnapshotRepository
from app.schemas.schedule import ScheduleWindow
from app.services.schedule_bindings import resolve_schedule_parameters
from app.services.snapshot_filtering import (
compile_snapshot_filters,
validate_snapshot_rows_match_resolved_parameters,
)
from app.sql.executor import execute_query

log = structlog.get_logger().bind(
Expand Down Expand Up @@ -236,6 +240,12 @@ async def execute_scheduled_job(
mapped_rows.append(new_row)
rows = mapped_rows

validate_snapshot_rows_match_resolved_parameters(
rows=rows,
filters=compile_snapshot_filters(param_schema),
resolved_params=params,
)

# Save snapshot
snapshot = Snapshot(
endpoint_id=eid,
Expand Down
52 changes: 50 additions & 2 deletions backend/app/services/snapshot_filtering.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
"""Typed request filtering for persisted snapshot rows."""

import json
from collections.abc import Mapping
from dataclasses import dataclass
from datetime import datetime
from functools import lru_cache
Expand Down Expand Up @@ -96,8 +97,8 @@ def _coerce_cached_row_value(item: CompiledSnapshotFilter, value: object) -> Any
def snapshot_covers_request(
*,
filters: tuple[CompiledSnapshotFilter, ...],
request_params: dict[str, object],
resolved_params: dict[str, object],
request_params: Mapping[str, object],
resolved_params: Mapping[str, object],
) -> bool:
"""Return whether one snapshot job run contains the requested selection."""
for item in filters:
Expand Down Expand Up @@ -174,6 +175,9 @@ def unavailable_snapshot_filter_columns(
"""Return configured output columns absent from a non-empty snapshot."""
if not rows:
return []
# TODO: Resolve configured filter columns against snapshot keys case-insensitively before both
# availability validation and row filtering, while rejecting ambiguous keys that differ only
# by case.
available = {column for row in rows for column in row}
return sorted({item.column for item in filters} - available)

Expand Down Expand Up @@ -219,3 +223,47 @@ def filter_snapshot_rows(
if matches:
filtered.append(row)
return filtered


def validate_snapshot_rows_match_resolved_parameters(
*,
rows: list[dict[str, object]],
filters: tuple[CompiledSnapshotFilter, ...],
resolved_params: Mapping[str, object],
) -> None:
"""Reject non-empty snapshot results that contradict their resolved schedule bounds."""
if not rows:
return

missing_columns = unavailable_snapshot_filter_columns(rows=rows, filters=filters)
if missing_columns:
raise ValueError(
"Snapshot integrity validation failed: configured filter columns are absent from "
f"the cached output: {', '.join(missing_columns)}."
)

normalized_params = dict(resolved_params)
for item in filters:
if item.parameter not in resolved_params or resolved_params[item.parameter] is None:
continue
try:
normalized_params[item.parameter] = item.coerce_value(resolved_params[item.parameter])
except (ValidationError, ValueError, TypeError) as exc:
raise ValueError(
"Snapshot integrity validation failed: the schedule's resolved parameter "
f":{item.parameter} cannot be coerced to {item.param_type}."
) from exc

matching_rows = filter_snapshot_rows(
rows=rows,
filters=filters,
request_params=normalized_params,
)
invalid_row_count = len(rows) - len(matching_rows)
if invalid_row_count:
columns = sorted({item.column for item in filters})
raise ValueError(
"Snapshot integrity validation failed: "
f"{invalid_row_count} of {len(rows)} cached rows do not match the schedule's "
f"resolved filter parameters for columns: {', '.join(columns)}."
)
Loading