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
115 changes: 55 additions & 60 deletions backend/secuscan/database.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,8 +44,6 @@ async def connect(self):
conn.row_factory = aiosqlite.Row
await conn.execute("PRAGMA foreign_keys = ON")
await self._create_schema()
await self._ensure_schema_migrations_table()
await self._validate_schema_version()
await self._run_migrations()

async def disconnect(self):
Expand Down Expand Up @@ -658,6 +656,60 @@ async def _create_schema(self):
except Exception as e:
print(f"Failed to add 'schedule_timezone' to workflows: {e}")

# Saved views table migration: ensure owner_id and composite unique exist
try:
saved_views_columns = await self.fetchall("PRAGMA table_info(saved_views)")
if saved_views_columns:
existing_sv_cols = {col["name"] for col in saved_views_columns}
if "owner_id" not in existing_sv_cols:
try:
await self.execute(
"ALTER TABLE saved_views ADD COLUMN owner_id TEXT NOT NULL DEFAULT 'default'"
)
existing_sv_cols.add("owner_id")
print("Added missing column 'owner_id' to saved_views table.")
except Exception as e:
print(f"Failed to add 'owner_id' to saved_views: {e}")

sv_schema = await self.fetchone(
"SELECT sql FROM sqlite_master WHERE type='table' AND name='saved_views'"
)
if sv_schema and "owner_id" in existing_sv_cols:
ddl = sv_schema["sql"]
has_old_unique = "name TEXT NOT NULL UNIQUE" in ddl
has_composite = "UNIQUE(owner_id, name)" in ddl
if has_old_unique or not has_composite:
old_fk = await self.fetchone("PRAGMA foreign_keys")
if old_fk:
await self.execute("PRAGMA foreign_keys = OFF")
try:
await self.connection.executescript("""
CREATE TABLE saved_views_new (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
owner_id TEXT NOT NULL DEFAULT 'default',
filter_json TEXT NOT NULL,
created_at TIMESTAMP NOT NULL DEFAULT (datetime('now')),
updated_at TIMESTAMP NOT NULL DEFAULT (datetime('now')),
UNIQUE(owner_id, name)
);
INSERT INTO saved_views_new
(id, name, owner_id, filter_json, created_at, updated_at)
SELECT
id, name, COALESCE(owner_id, 'default'),
filter_json, created_at, updated_at
FROM saved_views;
DROP TABLE saved_views;
ALTER TABLE saved_views_new RENAME TO saved_views;
""")
await self.connection.commit()
print("Replaced saved_views UNIQUE(name) constraint with UNIQUE(owner_id, name).")
finally:
if old_fk:
await self.execute("PRAGMA foreign_keys = ON")
except sqlite3.OperationalError:
pass # Table may not exist yet if migrations haven't run

# Notification rules table migration: ensure owner_id exists
notif_columns = await self.fetchall("PRAGMA table_info(notification_rules)")
existing_notif_cols = {col["name"] for col in notif_columns}
Expand Down Expand Up @@ -703,54 +755,6 @@ async def _create_schema(self):
)


async def _ensure_schema_migrations_table(self):
"""Create the migration tracking table if it does not already exist."""
await self.connection.execute(
"""
CREATE TABLE IF NOT EXISTS schema_migrations (
version TEXT PRIMARY KEY,
applied_at TIMESTAMP NOT NULL DEFAULT (datetime('now'))
)
"""
)
await self.connection.commit()

async def _applied_migrations(self) -> set[str]:
"""Return the set of migration filenames already applied."""
rows = await self.fetchall(
"SELECT version FROM schema_migrations"
)
return {row["version"] for row in rows}

async def _validate_schema_version(self):
"""Ensure the database was not created by a newer application."""

applied = await self._applied_migrations()

available = {
migration.name
for migration in (Path(__file__).parent / "migrations").glob("*.sql")
}

unknown = applied - available

if unknown:
raise RuntimeError(
"Database schema is newer than this application. "
f"Unknown migration(s): {', '.join(sorted(unknown))}"
)

async def _record_migration(self, version: str):
"""Record a successfully applied migration."""
await self.execute(
"""
INSERT INTO schema_migrations(version)
VALUES (?)
""",
(version,),
)


async def _run_migrations(self):
migrations_dir = Path(__file__).parent / "migrations"

Expand All @@ -760,22 +764,13 @@ async def _run_migrations(self):
"ensure the backend package is installed correctly."
)

applied = await self._applied_migrations()

for migration_file in sorted(migrations_dir.glob("*.sql")):
migration_name = migration_file.name

if migration_name in applied:
continue

sql = migration_file.read_text(encoding="utf-8")

try:
await self.connection.executescript(sql)
await self._record_migration(migration_name)
except Exception as exc:
raise RuntimeError(
f"Migration {migration_name} failed — startup aborted: {exc}"
f"Migration {migration_file.name} failed — startup aborted: {exc}"
) from exc

await self._backfill_risk_scores()
Expand Down
7 changes: 5 additions & 2 deletions backend/secuscan/migrations/002_add_saved_views.sql
Original file line number Diff line number Diff line change
Expand Up @@ -12,10 +12,13 @@

CREATE TABLE IF NOT EXISTS saved_views (
id TEXT PRIMARY KEY,
name TEXT NOT NULL UNIQUE,
name TEXT NOT NULL,
owner_id TEXT NOT NULL DEFAULT 'default',
filter_json TEXT NOT NULL,
created_at TIMESTAMP NOT NULL DEFAULT (datetime('now')),
updated_at TIMESTAMP NOT NULL DEFAULT (datetime('now'))
updated_at TIMESTAMP NOT NULL DEFAULT (datetime('now')),
UNIQUE(owner_id, name)
);

CREATE INDEX IF NOT EXISTS idx_saved_views_name ON saved_views(LOWER(name));
CREATE INDEX IF NOT EXISTS idx_saved_views_owner ON saved_views(owner_id);
41 changes: 19 additions & 22 deletions backend/secuscan/saved_views.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,17 +4,13 @@
import uuid
from typing import Any, Dict, List, Optional

from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, HTTPException, Depends
from pydantic import BaseModel, Field, field_validator

from .auth import require_api_key
from .database import get_db
from .auth import get_current_owner

saved_views_router = APIRouter(
prefix="/api/v1/saved-views",
tags=["saved-views"],
dependencies=[Depends(require_api_key)],
)
saved_views_router = APIRouter(prefix="/api/v1/saved-views", tags=["saved-views"])

_VALID_SORT_MODES = {"severity", "newest", "oldest", "target"}
_VALID_SEVERITIES = {"all", "critical", "high", "medium", "low", "info"}
Expand Down Expand Up @@ -101,18 +97,19 @@ def validate_filter_json(cls, v: Optional[str]) -> Optional[str]:


@saved_views_router.get("")
async def list_saved_views() -> Dict[str, Any]:
async def list_saved_views(owner_id: str = Depends(get_current_owner)) -> Dict[str, Any]:
"""Return all saved views ordered by creation date."""
db = await get_db()
rows: List[Dict] = await db.fetchall(
"SELECT id, name, filter_json, created_at, updated_at "
"FROM saved_views ORDER BY created_at ASC"
"FROM saved_views WHERE owner_id = ? ORDER BY created_at ASC",
(owner_id,)
)
return {"views": rows, "total": len(rows)}


@saved_views_router.post("", status_code=201)
async def create_saved_view(body: SavedViewCreate) -> Dict[str, Any]:
async def create_saved_view(body: SavedViewCreate, owner_id: str = Depends(get_current_owner)) -> Dict[str, Any]:
"""
Create a new saved view.
Returns 409 if a view with the same name already exists.
Expand All @@ -121,7 +118,7 @@ async def create_saved_view(body: SavedViewCreate) -> Dict[str, Any]:


existing = await db.fetchone(
"SELECT id FROM saved_views WHERE LOWER(name) = LOWER(?)", (body.name,)
"SELECT id FROM saved_views WHERE owner_id = ? AND LOWER(name) = LOWER(?)", (owner_id, body.name)
)
if existing:
raise HTTPException(
Expand All @@ -133,23 +130,23 @@ async def create_saved_view(body: SavedViewCreate) -> Dict[str, Any]:
view_id = str(uuid.uuid4())
await db.execute(
"""
INSERT INTO saved_views (id, name, filter_json)
VALUES (?, ?, ?)
INSERT INTO saved_views (id, name, owner_id, filter_json)
VALUES (?, ?, ?, ?)
""",
(view_id, body.name, body.filter_json),
(view_id, body.name, owner_id, body.filter_json),
)
return {"id": view_id, "name": body.name, "created": True}


@saved_views_router.put("/{view_id}")
async def update_saved_view(view_id: str, body: SavedViewUpdate) -> Dict[str, Any]:
async def update_saved_view(view_id: str, body: SavedViewUpdate, owner_id: str = Depends(get_current_owner)) -> Dict[str, Any]:
"""
Overwrite name and/or filter_json for an existing view.
Also accepts PATCH semantics — only supplied fields are updated.
"""
db = await get_db()

row = await db.fetchone("SELECT id FROM saved_views WHERE id = ?", (view_id,))
row = await db.fetchone("SELECT id FROM saved_views WHERE id = ? AND owner_id = ?", (view_id, owner_id))
if not row:
raise HTTPException(status_code=404, detail="Saved view not found")

Expand All @@ -159,8 +156,8 @@ async def update_saved_view(view_id: str, body: SavedViewUpdate) -> Dict[str, An
if body.name is not None:
# Check for name collision with a *different* record
collision = await db.fetchone(
"SELECT id FROM saved_views WHERE LOWER(name) = LOWER(?) AND id != ?",
(body.name, view_id),
"SELECT id FROM saved_views WHERE owner_id = ? AND LOWER(name) = LOWER(?) AND id != ?",
(owner_id, body.name, view_id),
)
if collision:
raise HTTPException(
Expand All @@ -181,15 +178,15 @@ async def update_saved_view(view_id: str, body: SavedViewUpdate) -> Dict[str, An
params.append(view_id)

await db.execute(
f"UPDATE saved_views SET {', '.join(updates)} WHERE id = ?",
tuple(params),
f"UPDATE saved_views SET {', '.join(updates)} WHERE id = ? AND owner_id = ?",
tuple(params) + (owner_id,),
)
return {"id": view_id, "updated": True}


@saved_views_router.delete("/{view_id}")
async def delete_saved_view(view_id: str) -> Dict[str, Any]:
async def delete_saved_view(view_id: str, owner_id: str = Depends(get_current_owner)) -> Dict[str, Any]:
"""Delete a saved view by id. Idempotent — returns 200 even if not found."""
db = await get_db()
await db.execute("DELETE FROM saved_views WHERE id = ?", (view_id,))
await db.execute("DELETE FROM saved_views WHERE id = ? AND owner_id = ?", (view_id, owner_id))
return {"id": view_id, "deleted": True}
Loading