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
13 changes: 12 additions & 1 deletion kubetix-api/kubetix_api/cleanup.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
from sqlalchemy.orm import Session

from kubetix_api.database import SessionLocal
from kubetix_api.models import Grant
from kubetix_api.models import AuditLog, Grant

log = logging.getLogger(__name__)

Expand Down Expand Up @@ -47,12 +47,23 @@ async def run_grant_cleanup_loop(stop_event: asyncio.Event) -> None:
def purge_expired_grants(session_factory=SessionLocal) -> int:
"""Delete grants whose ``expires_at`` is in the past.

Before deleting, null out grant_id on any AuditLog rows that reference the
expired grants so the forensic audit trail is preserved (soft-reference).

Returns the number of rows removed. ``session_factory`` is a callable
that returns a SQLAlchemy ``Session`` (defaults to ``SessionLocal``).
"""
db: Session = session_factory()
try:
now = datetime.now(timezone.utc)
expired_grants = db.query(Grant).filter(Grant.expires_at < now).all()
if not expired_grants:
return 0
expired_ids = [g.id for g in expired_grants]
# Preserve audit trail: null out grant_id references before deleting grants
db.query(AuditLog).filter(AuditLog.grant_id.in_(expired_ids)).update(
{AuditLog.grant_id: None}, synchronize_session="fetch"
)
deleted = (
db.query(Grant)
.filter(Grant.expires_at < now)
Expand Down
35 changes: 31 additions & 4 deletions kubetix-api/tests/test_cleanup.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,22 @@
from unittest.mock import MagicMock

from kubetix_api.cleanup import purge_expired_grants
from kubetix_api.models import Grant
from kubetix_api.models import AuditLog, Grant


def _make_session_factory(expired_count):
"""Return a session factory whose Grant.query.delete returns ``expired_count``."""
session = MagicMock()
# When there are expired grants, .all() returns mock grant objects with ids
if expired_count > 0:
fake_grants = [MagicMock(id=f"grant-{i}") for i in range(expired_count)]
session.query.return_value.filter.return_value.all.return_value = fake_grants
# AuditLog query update returns the number of rows updated
session.query.return_value.filter.return_value.update.return_value = (
expired_count
)
else:
session.query.return_value.filter.return_value.all.return_value = []
session.query.return_value.filter.return_value.delete.return_value = expired_count
factory = MagicMock(return_value=session)
return factory, session
Expand All @@ -21,8 +31,8 @@ def test_purge_deletes_expired_grants_and_commits():

assert deleted == 3
# delete() is called on a Query object filtered by expires_at < now
session.query.assert_called_once_with(Grant)
session.query.return_value.filter.assert_called_once()
session.query.assert_called()
session.query.return_value.filter.assert_called()
session.commit.assert_called_once()
session.close.assert_called_once()

Expand All @@ -33,7 +43,7 @@ def test_purge_returns_zero_when_no_expired_grants():
deleted = purge_expired_grants(factory)

assert deleted == 0
session.commit.assert_called_once()
session.commit.assert_not_called()
session.close.assert_called_once()


Expand All @@ -48,3 +58,20 @@ def test_purge_rolls_back_on_error():

session.rollback.assert_called_once()
session.close.assert_called_once()


def test_purge_nulls_audit_log_grant_id_before_delete():
"""AuditLog grant_id is nulled before expired grants are deleted (issue #250)."""
factory, session = _make_session_factory(expired_count=2)

purge_expired_grants(factory)

# Verify AuditLog was queried and updated with grant_id=None
query_calls = session.query.call_args_list
assert any(
call[0][0].__name__ == "AuditLog" for call in query_calls
), "AuditLog should be queried before deleting grants"
# The update call should set grant_id to None
session.query.return_value.filter.return_value.update.assert_called_once()
update_call = session.query.return_value.filter.return_value.update.call_args
assert update_call[0][0] == {AuditLog.grant_id: None}
103 changes: 95 additions & 8 deletions tests/unit/test_cleanup.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,11 @@
def test_purge_expired_grants_removes_only_expired(monkeypatch):
"""Expired grants are deleted; non-expired grants are kept."""

class FakeGrant:
def __init__(self, id_, expires_at):
self.id = id_
self.expires_at = expires_at

class FakeSession:
def __init__(self, store):
self.store = store
Expand All @@ -18,11 +23,23 @@ def query(self, model):

def filter(self, _cond):
now = datetime.now(timezone.utc)
kept = [g for g in self.store if g["expires_at"] >= now]
removed = len(self.store) - len(kept)
expired = [g for g in self.store if g.expires_at < now]
kept = [g for g in self.store if g.expires_at >= now]
self._expired = expired
self._kept = kept
return self

def all(self):
return list(getattr(self, "_expired", []))

def update(self, values, synchronize_session="fetch"):
return 0

def delete(self, synchronize_session=False):
removed = len(self._expired)
self.store.clear()
self.store.extend(kept)
return MagicMock(delete=MagicMock(return_value=removed))
self.store.extend(self._kept)
return removed

def commit(self):
pass
Expand All @@ -35,9 +52,9 @@ def close(self):

now = datetime.now(timezone.utc)
store = [
{"id": 1, "expires_at": now - timedelta(hours=1)}, # expired
{"id": 2, "expires_at": now + timedelta(hours=1)}, # active
{"id": 3, "expires_at": now - timedelta(seconds=1)}, # expired
FakeGrant(1, now - timedelta(hours=1)), # expired
FakeGrant(2, now + timedelta(hours=1)), # active
FakeGrant(3, now - timedelta(seconds=1)), # expired
]

def factory():
Expand All @@ -48,7 +65,7 @@ def factory():
deleted = cleanup.purge_expired_grants(factory)

assert deleted == 2
assert [g["id"] for g in store] == [2]
assert [g.id for g in store] == [2]


@pytest.mark.asyncio
Expand Down Expand Up @@ -108,3 +125,73 @@ def flaky_purge(_factory):
assert any(
"Expired-grant cleanup iteration failed" in r.message for r in caplog.records
)


def test_purge_expired_grants_preserves_audit_trail(monkeypatch):
"""AuditLog rows referencing expired grants have grant_id nulled before deletion (issue #250)."""

class FakeGrant:
def __init__(self, id_):
self.id = id_

class FakeAuditLog:
def __init__(self, grant_id):
self.grant_id = grant_id

class FakeQuery:
def __init__(self, model, store):
self.model = model
self.store = store

def filter(self, _cond):
return self

def all(self):
return list(self.store)

def update(self, values, synchronize_session="fetch"):
for item in self.store:
for key, val in values.items():
# Handle SQLAlchemy InstrumentedAttribute keys
attr_name = getattr(key, "key", str(key))
setattr(item, attr_name, val)
return len(self.store)

def delete(self, synchronize_session=False):
count = len(self.store)
self.store.clear()
return count

class FakeSession:
def __init__(self):
self.grants = [
FakeGrant("expired-1"),
FakeGrant("expired-2"),
]
self.audit_logs = [
FakeAuditLog("expired-1"),
FakeAuditLog("expired-2"),
FakeAuditLog(None),
]

def query(self, model):
if model.__name__ == "Grant":
return FakeQuery(model, self.grants)
elif model.__name__ == "AuditLog":
return FakeQuery(model, self.audit_logs)
raise ValueError(f"Unknown model: {model}")

def commit(self):
pass

def rollback(self):
raise RuntimeError("rollback called unexpectedly")

def close(self):
pass

from kubetix_api import cleanup

deleted = cleanup.purge_expired_grants(lambda: FakeSession())

assert deleted == 2