Skip to content
Draft
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
44 changes: 43 additions & 1 deletion src/app/api/api_v1/endpoints/sequences.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
from app.schemas.alerts import AlertCreate
from app.schemas.detections import DetectionSequence, DetectionWithUrl
from app.schemas.login import TokenPayload
from app.schemas.sequences import SequenceLabel, SequenceRead
from app.schemas.sequences import SequenceAzimuth, SequenceLabel, SequenceRead
from app.services.alerts import refresh_alert_state
from app.services.risk import FwiClass, risk_service
from app.services.sequence_confidence import max_conf_filter_clause
Expand Down Expand Up @@ -323,3 +323,45 @@ async def label_sequence(
await session.commit()

return updated


@router.patch(
"/{sequence_id}/azimuth", status_code=status.HTTP_200_OK, summary="Refine the azimuth of the sequence"
)
async def refine_azimuth(
payload: SequenceAzimuth,
sequence_id: int = Path(..., gt=0),
cameras: CameraCRUD = Depends(get_camera_crud),
sequences: SequenceCRUD = Depends(get_sequence_crud),
alerts: AlertCRUD = Depends(get_alert_crud),
session: AsyncSession = Depends(get_session),
token_payload: TokenPayload = Security(get_jwt, scopes=[UserRole.ADMIN, UserRole.USER]),
) -> Sequence:

telemetry_client.capture(token_payload.sub, event="azimuth-refine", properties={"sequence_id": sequence_id})

# Fetch the sequence:
sequence = cast(Sequence, await sequences.get(sequence_id, strict=True))

# Non-admins are scoped to their own organization:
if not token_payload.is_admin:
await verify_org_rights(token_payload.organization_id, sequence.camera_id, cameras)

# Persist the new azimuth:
updated = await sequences.update(sequence_id, payload)

# Re-run the same attach/reconcile flow:
# Imported lazily to avoid a services -> endpoints import at module load.
from app.api.api_v1.endpoints.detections import _attach_sequence_to_alert

camera = cast(Camera, await cameras.get(sequence.camera_id, strict=True))
# Recompute triangulation:
_ = _attach_sequence_to_alert(sequence_=updated, camera=camera, cameras=cameras, sequences=sequences, alerts=alerts)

# Refresh every previously- and newly-linked alert:
alert_ids_res = await session.exec(select(AlertSequence.alert_id).where(AlertSequence.sequence_id == sequence_id))
alert_ids = list(alert_ids_res.all())
for aid in alert_ids:
await refresh_alert_state(aid, session, alerts)

return updated
4 changes: 2 additions & 2 deletions src/app/crud/crud_sequence.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,14 +15,14 @@
from app.core.time import utcnow
from app.crud.base import BaseCRUD
from app.models import TERMINAL_VALIDATION_STATUSES, VALIDATION_FAILED, Detection, Sequence
from app.schemas.sequences import SequenceLabel, SequenceUpdate
from app.schemas.sequences import SequenceAzimuth, SequenceLabel, SequenceUpdate

__all__ = ["SequenceCRUD"]

logger = logging.getLogger("uvicorn.error")


class SequenceCRUD(BaseCRUD[Sequence, Sequence, Union[SequenceUpdate, SequenceLabel]]):
class SequenceCRUD(BaseCRUD[Sequence, Sequence, Union[SequenceAzimuth, SequenceUpdate, SequenceLabel]]):
def __init__(self, session: AsyncSession) -> None:
super().__init__(session, Sequence)

Expand Down
6 changes: 5 additions & 1 deletion src/app/schemas/sequences.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@

from app.models import AnnotationType, Sequence

__all__ = ["SequenceLabel", "SequenceRead", "SequenceUpdate"]
__all__ = ["SequenceAzimuth", "SequenceLabel", "SequenceRead", "SequenceUpdate"]


# Accesses
Expand All @@ -30,3 +30,7 @@ class SequenceRead(Sequence):
validation_lease_until: Union[datetime, None] = Field(None, exclude=True)
validation_status: Union[str, None] = Field(None, exclude=True)
validation_attempts: int = Field(0, exclude=True)


class SequenceAzimuth(BaseModel):
sequence_azimuth: float = Field(..., ge=0, lt=360)
33 changes: 33 additions & 0 deletions src/tests/endpoints/test_sequences.py
Original file line number Diff line number Diff line change
Expand Up @@ -224,6 +224,39 @@ async def test_label_sequence(
**payload,
}

@pytest.mark.parametrize(
("user_idx", "sequence_id", "payload", "status_code", "status_detail"),
[
(None, 1, {"sequence_azimuth": 100}, 401, "Not authenticated"), # No user auth
(0, 0, {"sequence_azimuth": 100}, 422, None), # Sequence id should be >0
(0, 0, {"sequence_azimuth": 700}, 422, None), # Sequence azimuth should be in range [0, 360)
(0, 99, {"sequence_azimuth": 100}, 404, None), # Inexsisting Sequence
(1, 1, {"sequence_azimuth": 100}, 403, None), # Agent does not have permission
(2, 1, {"sequence_azimuth": 100}, 403, "Access forbidden."), # Non admin action on out of scope organisation
],
)
@pytest.mark.asyncio
async def test_refine_azimuth(
async_client: AsyncClient,
sequence_session: AsyncSession,
user_idx: Union[int, None],
sequence_id: int,
payload: Dict[str, Any],
status_code: int,
status_detail: Union[str, None],
):
auth = None
if isinstance(user_idx, int):
auth = pytest.get_token(
pytest.user_table[user_idx]["id"],
pytest.user_table[user_idx]["role"].split(),
pytest.user_table[user_idx]["organization_id"],
)

response = await async_client.patch(f"/sequences/{sequence_id}/azimuth", json=payload, headers=auth)
assert response.status_code == status_code, print(response.__dict__)
if isinstance(status_detail, str):
assert response.json()["detail"] == status_detail

@pytest.mark.parametrize(
("user_idx", "from_date", "status_code", "status_detail", "expected_result"),
Expand Down
Loading