Harden production workflows and agent admission
This commit is contained in:
@@ -1,49 +1,341 @@
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy import and_, delete, or_, select
|
||||
from sqlalchemy.dialects.postgresql import insert as pg_insert
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ..cursors import decode, encode
|
||||
from ..db import get_session
|
||||
from ..models import Agent as AgentModel
|
||||
from ..schemas import Agent, AgentsPage
|
||||
from ..models import AgentBlacklist, Event, Incident, NotificationOutbox
|
||||
from ..schemas import (
|
||||
ActionCountResponse,
|
||||
Agent,
|
||||
AgentBlacklistItem,
|
||||
AgentBlacklistPage,
|
||||
AgentsPage,
|
||||
)
|
||||
from ..timeutil import to_utc
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _is_pending(a: AgentModel) -> bool:
|
||||
return a.pending_key_hash is not None
|
||||
|
||||
|
||||
def _to_schema(a: AgentModel) -> Agent:
|
||||
pending = _is_pending(a)
|
||||
hostname = a.pending_hostname if pending and a.pending_hostname is not None else a.hostname
|
||||
last_seen_at = (
|
||||
a.pending_seen_at if pending and a.pending_seen_at is not None else a.last_seen_at
|
||||
)
|
||||
features = a.pending_features if pending and a.pending_features is not None else a.features
|
||||
labels = a.pending_labels if pending and a.pending_labels is not None else a.labels
|
||||
return Agent(
|
||||
agent_id=a.agent_id,
|
||||
hostname=a.hostname,
|
||||
hostname=hostname,
|
||||
status=a.status,
|
||||
last_seen_at=to_utc(a.last_seen_at),
|
||||
features=a.features or {"push": False, "metrics": False},
|
||||
labels=a.labels or {},
|
||||
admission_state="pending" if pending or a.accepted_key_hash is None else "accepted",
|
||||
rotation_pending=pending and a.accepted_key_hash is not None,
|
||||
key_fingerprint=(a.pending_key_fingerprint if pending else a.accepted_key_fingerprint),
|
||||
last_seen_at=to_utc(last_seen_at),
|
||||
features=features or {"push": False, "metrics": False},
|
||||
labels=labels or {},
|
||||
)
|
||||
|
||||
|
||||
def _blacklist_schema(row: AgentBlacklist) -> AgentBlacklistItem:
|
||||
return AgentBlacklistItem(
|
||||
id=row.key_hash,
|
||||
key_fingerprint=row.key_fingerprint,
|
||||
agent_id=row.agent_id,
|
||||
hostname=row.hostname,
|
||||
created_at=to_utc(row.created_at),
|
||||
last_seen_at=to_utc(row.last_seen_at),
|
||||
)
|
||||
|
||||
|
||||
async def _get_agent(session: AsyncSession, agent_id: str) -> AgentModel:
|
||||
res = await session.execute(select(AgentModel).where(AgentModel.agent_id == agent_id))
|
||||
a = res.scalar_one_or_none()
|
||||
if a is None:
|
||||
raise HTTPException(status_code=404, detail="agent not found")
|
||||
return a
|
||||
|
||||
|
||||
def _clear_pending(a: AgentModel) -> None:
|
||||
a.pending_key_hash = None
|
||||
a.pending_key_fingerprint = None
|
||||
a.pending_hostname = None
|
||||
a.pending_seen_at = None
|
||||
a.pending_features = None
|
||||
a.pending_labels = None
|
||||
|
||||
|
||||
def _accept_pending(a: AgentModel) -> bool:
|
||||
if not _is_pending(a):
|
||||
return False
|
||||
now = datetime.now(UTC)
|
||||
a.accepted_key_hash = a.pending_key_hash
|
||||
a.accepted_key_fingerprint = a.pending_key_fingerprint
|
||||
a.accepted_at = now
|
||||
a.hostname = a.pending_hostname or a.hostname
|
||||
a.last_seen_at = a.pending_seen_at or a.last_seen_at
|
||||
if a.pending_features is not None:
|
||||
a.features = a.pending_features
|
||||
if a.pending_labels is not None:
|
||||
a.labels = a.pending_labels
|
||||
a.status = "alive"
|
||||
_clear_pending(a)
|
||||
return True
|
||||
|
||||
|
||||
async def _block_pending(session: AsyncSession, a: AgentModel) -> bool:
|
||||
if not _is_pending(a):
|
||||
return False
|
||||
now = datetime.now(UTC)
|
||||
stmt = pg_insert(AgentBlacklist).values(
|
||||
key_hash=a.pending_key_hash,
|
||||
key_fingerprint=a.pending_key_fingerprint,
|
||||
agent_id=a.agent_id,
|
||||
hostname=a.pending_hostname or a.hostname,
|
||||
last_seen_at=a.pending_seen_at or now,
|
||||
)
|
||||
stmt = stmt.on_conflict_do_update(
|
||||
index_elements=[AgentBlacklist.key_hash],
|
||||
set_={
|
||||
"agent_id": stmt.excluded.agent_id,
|
||||
"hostname": stmt.excluded.hostname,
|
||||
"last_seen_at": stmt.excluded.last_seen_at,
|
||||
},
|
||||
)
|
||||
await session.execute(stmt)
|
||||
_clear_pending(a)
|
||||
if a.accepted_key_hash is None:
|
||||
await _delete_agent_data(session, a.agent_id)
|
||||
return True
|
||||
|
||||
|
||||
async def _delete_agent_data(session: AsyncSession, agent_id: str) -> None:
|
||||
incident_ids = select(Incident.id).where(Incident.agent_id == agent_id)
|
||||
await session.execute(
|
||||
delete(NotificationOutbox).where(NotificationOutbox.incident_id.in_(incident_ids))
|
||||
)
|
||||
await session.execute(delete(Incident).where(Incident.agent_id == agent_id))
|
||||
await session.execute(delete(Event).where(Event.agent_id == agent_id))
|
||||
await session.execute(delete(AgentModel).where(AgentModel.agent_id == agent_id))
|
||||
|
||||
|
||||
async def _remove_pending_or_agent(session: AsyncSession, a: AgentModel) -> None:
|
||||
if _is_pending(a) and a.accepted_key_hash is not None:
|
||||
_clear_pending(a)
|
||||
return
|
||||
await _delete_agent_data(session, a.agent_id)
|
||||
|
||||
|
||||
def _agent_rank(a: AgentModel) -> int:
|
||||
return a.admission_rank
|
||||
|
||||
|
||||
def _decode_agent_cursor(cursor: str) -> tuple[int, str]:
|
||||
rank_raw, agent_id = decode(cursor, 2)
|
||||
try:
|
||||
rank = int(rank_raw)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail="invalid cursor") from exc
|
||||
if rank not in (0, 1):
|
||||
raise HTTPException(status_code=400, detail="invalid cursor")
|
||||
return rank, agent_id
|
||||
|
||||
|
||||
@router.get("/agents", response_model=AgentsPage)
|
||||
async def list_agents(
|
||||
cursor: str | None = Query(default=None, max_length=512),
|
||||
back: bool = Query(default=False),
|
||||
limit: int = Query(default=100, ge=1, le=500),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> AgentsPage:
|
||||
stmt = select(AgentModel).order_by(AgentModel.agent_id)
|
||||
if back:
|
||||
stmt = select(AgentModel).order_by(
|
||||
AgentModel.admission_rank.desc(), AgentModel.agent_id.desc()
|
||||
)
|
||||
else:
|
||||
stmt = select(AgentModel).order_by(
|
||||
AgentModel.admission_rank.asc(), AgentModel.agent_id.asc()
|
||||
)
|
||||
if cursor is not None:
|
||||
(last_id,) = decode(cursor, 1)
|
||||
stmt = stmt.where(AgentModel.agent_id > last_id)
|
||||
last_rank, last_id = _decode_agent_cursor(cursor)
|
||||
if back:
|
||||
stmt = stmt.where(
|
||||
or_(
|
||||
AgentModel.admission_rank < last_rank,
|
||||
and_(
|
||||
AgentModel.admission_rank == last_rank,
|
||||
AgentModel.agent_id < last_id,
|
||||
),
|
||||
)
|
||||
)
|
||||
else:
|
||||
stmt = stmt.where(
|
||||
or_(
|
||||
AgentModel.admission_rank > last_rank,
|
||||
and_(
|
||||
AgentModel.admission_rank == last_rank,
|
||||
AgentModel.agent_id > last_id,
|
||||
),
|
||||
)
|
||||
)
|
||||
rows = (await session.execute(stmt.limit(limit + 1))).scalars().all()
|
||||
next_cursor = None
|
||||
if len(rows) > limit:
|
||||
rows = rows[:limit]
|
||||
next_cursor = encode(rows[-1].agent_id)
|
||||
return AgentsPage(items=[_to_schema(a) for a in rows], next_cursor=next_cursor)
|
||||
has_more = len(rows) > limit
|
||||
rows = rows[:limit]
|
||||
if back:
|
||||
rows = list(reversed(rows))
|
||||
next_cursor: str | None = None
|
||||
prev_cursor: str | None = None
|
||||
if rows:
|
||||
if back:
|
||||
next_cursor = encode(_agent_rank(rows[-1]), rows[-1].agent_id)
|
||||
if has_more:
|
||||
prev_cursor = encode(_agent_rank(rows[0]), rows[0].agent_id)
|
||||
else:
|
||||
if has_more:
|
||||
next_cursor = encode(_agent_rank(rows[-1]), rows[-1].agent_id)
|
||||
if cursor is not None:
|
||||
prev_cursor = encode(_agent_rank(rows[0]), rows[0].agent_id)
|
||||
return AgentsPage(
|
||||
items=[_to_schema(a) for a in rows],
|
||||
next_cursor=next_cursor,
|
||||
prev_cursor=prev_cursor,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/agents/admission/blacklist", response_model=AgentBlacklistPage)
|
||||
async def list_agent_blacklist(
|
||||
cursor: str | None = Query(default=None, max_length=512),
|
||||
limit: int = Query(default=100, ge=1, le=500),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> AgentBlacklistPage:
|
||||
stmt = select(AgentBlacklist).order_by(AgentBlacklist.key_fingerprint, AgentBlacklist.key_hash)
|
||||
if cursor is not None:
|
||||
last_fingerprint, last_hash = decode(cursor, 2)
|
||||
stmt = stmt.where(
|
||||
or_(
|
||||
AgentBlacklist.key_fingerprint > last_fingerprint,
|
||||
and_(
|
||||
AgentBlacklist.key_fingerprint == last_fingerprint,
|
||||
AgentBlacklist.key_hash > last_hash,
|
||||
),
|
||||
)
|
||||
)
|
||||
rows = (await session.execute(stmt.limit(limit + 1))).scalars().all()
|
||||
has_more = len(rows) > limit
|
||||
rows = rows[:limit]
|
||||
return AgentBlacklistPage(
|
||||
items=[_blacklist_schema(r) for r in rows],
|
||||
next_cursor=encode(rows[-1].key_fingerprint, rows[-1].key_hash)
|
||||
if rows and has_more
|
||||
else None,
|
||||
prev_cursor=None,
|
||||
)
|
||||
|
||||
|
||||
@router.delete("/agents/admission/blacklist", response_model=ActionCountResponse)
|
||||
async def clear_agent_blacklist(
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> ActionCountResponse:
|
||||
res = await session.execute(delete(AgentBlacklist))
|
||||
await session.commit()
|
||||
return ActionCountResponse(count=res.rowcount or 0)
|
||||
|
||||
|
||||
@router.delete("/agents/admission/blacklist/{key_id}", response_model=ActionCountResponse)
|
||||
async def delete_agent_blacklist_item(
|
||||
key_id: str,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> ActionCountResponse:
|
||||
res = await session.execute(delete(AgentBlacklist).where(AgentBlacklist.key_hash == key_id))
|
||||
await session.commit()
|
||||
return ActionCountResponse(count=res.rowcount or 0)
|
||||
|
||||
|
||||
@router.post("/agents/admission/accept-all", response_model=ActionCountResponse)
|
||||
async def accept_all_pending_agents(
|
||||
include_rotation: bool = Query(default=False),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> ActionCountResponse:
|
||||
stmt = select(AgentModel).where(AgentModel.pending_key_hash.is_not(None))
|
||||
if not include_rotation:
|
||||
stmt = stmt.where(AgentModel.accepted_key_hash.is_(None))
|
||||
rows = (await session.execute(stmt)).scalars().all()
|
||||
count = sum(1 for a in rows if _accept_pending(a))
|
||||
await session.commit()
|
||||
return ActionCountResponse(count=count)
|
||||
|
||||
|
||||
@router.post("/agents/admission/block-all", response_model=ActionCountResponse)
|
||||
async def block_all_pending_agents(
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> ActionCountResponse:
|
||||
rows = (
|
||||
(await session.execute(select(AgentModel).where(AgentModel.pending_key_hash.is_not(None))))
|
||||
.scalars()
|
||||
.all()
|
||||
)
|
||||
count = 0
|
||||
for a in rows:
|
||||
if await _block_pending(session, a):
|
||||
count += 1
|
||||
await session.commit()
|
||||
return ActionCountResponse(count=count)
|
||||
|
||||
|
||||
@router.post("/agents/admission/remove-all", response_model=ActionCountResponse)
|
||||
async def remove_all_pending_agents(
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> ActionCountResponse:
|
||||
rows = (
|
||||
(await session.execute(select(AgentModel).where(AgentModel.pending_key_hash.is_not(None))))
|
||||
.scalars()
|
||||
.all()
|
||||
)
|
||||
for agent in rows:
|
||||
await _remove_pending_or_agent(session, agent)
|
||||
await session.commit()
|
||||
return ActionCountResponse(count=len(rows))
|
||||
|
||||
|
||||
@router.post("/agents/{agent_id}/accept", response_model=ActionCountResponse)
|
||||
async def accept_agent(
|
||||
agent_id: str, session: AsyncSession = Depends(get_session)
|
||||
) -> ActionCountResponse:
|
||||
agent = await _get_agent(session, agent_id)
|
||||
count = 1 if _accept_pending(agent) else 0
|
||||
await session.commit()
|
||||
return ActionCountResponse(count=count)
|
||||
|
||||
|
||||
@router.post("/agents/{agent_id}/block", response_model=ActionCountResponse)
|
||||
async def block_agent(
|
||||
agent_id: str, session: AsyncSession = Depends(get_session)
|
||||
) -> ActionCountResponse:
|
||||
agent = await _get_agent(session, agent_id)
|
||||
count = 1 if await _block_pending(session, agent) else 0
|
||||
await session.commit()
|
||||
return ActionCountResponse(count=count)
|
||||
|
||||
|
||||
@router.delete("/agents/{agent_id}", response_model=ActionCountResponse)
|
||||
async def delete_agent(
|
||||
agent_id: str, session: AsyncSession = Depends(get_session)
|
||||
) -> ActionCountResponse:
|
||||
agent = await _get_agent(session, agent_id)
|
||||
await _remove_pending_or_agent(session, agent)
|
||||
await session.commit()
|
||||
return ActionCountResponse(count=1)
|
||||
|
||||
|
||||
@router.get("/agents/{agent_id}", response_model=Agent)
|
||||
async def get_agent(agent_id: str, session: AsyncSession = Depends(get_session)) -> Agent:
|
||||
res = await session.execute(select(AgentModel).where(AgentModel.agent_id == agent_id))
|
||||
a = res.scalar_one_or_none()
|
||||
if a is None:
|
||||
raise HTTPException(status_code=404, detail="agent not found")
|
||||
return _to_schema(a)
|
||||
return _to_schema(await _get_agent(session, agent_id))
|
||||
|
||||
Reference in New Issue
Block a user