| 1 | from typing import List |
| 2 | |
| 3 | from fastapi import HTTPException |
| 4 | from loguru import logger |
| 5 | from sqlalchemy import select |
| 6 | from sqlalchemy.ext.asyncio import AsyncSession |
| 7 | |
| 8 | from app.db.universal_models import EventSources |
| 9 | from app.siem.schema.event_sources import EventSourceCreate |
| 10 | from app.siem.schema.event_sources import EventSourceUpdate |
| 11 | |
| 12 | |
| 13 | async def get_event_sources_by_customer( |
| 14 | customer_code: str, |
| 15 | db: AsyncSession, |
| 16 | ) -> List[EventSources]: |
| 17 | logger.info(f"Fetching event sources for customer {customer_code}") |
| 18 | result = await db.execute( |
| 19 | select(EventSources).where(EventSources.customer_code == customer_code), |
| 20 | ) |
| 21 | return result.scalars().all() |
| 22 | |
| 23 | |
| 24 | async def get_event_source_by_id( |
| 25 | event_source_id: int, |
| 26 | db: AsyncSession, |
| 27 | ) -> EventSources: |
| 28 | result = await db.execute( |
| 29 | select(EventSources).where(EventSources.id == event_source_id), |
| 30 | ) |
| 31 | event_source = result.scalars().first() |
| 32 | if not event_source: |
| 33 | raise HTTPException(status_code=404, detail="Event source not found") |
| 34 | return event_source |
| 35 | |
| 36 | |
| 37 | async def create_event_source( |
| 38 | event_source_data: EventSourceCreate, |
| 39 | db: AsyncSession, |
| 40 | ) -> EventSources: |
| 41 | logger.info(f"Creating event source '{event_source_data.name}' for customer {event_source_data.customer_code}") |
| 42 | # Check for duplicate name within the same customer |
| 43 | result = await db.execute( |
| 44 | select(EventSources).where( |
| 45 | EventSources.customer_code == event_source_data.customer_code, |
| 46 | EventSources.name == event_source_data.name, |
| 47 | ), |
| 48 | ) |
| 49 | if result.scalars().first(): |
| 50 | raise HTTPException( |
| 51 | status_code=400, |
| 52 | detail=f"Event source '{event_source_data.name}' already exists for customer {event_source_data.customer_code}", |
| 53 | ) |
| 54 | |
| 55 | db_event_source = EventSources(**event_source_data.model_dump()) |
| 56 | db.add(db_event_source) |
| 57 | await db.flush() |
| 58 | await db.refresh(db_event_source) |
| 59 | await db.commit() |
| 60 | return db_event_source |
| 61 | |
| 62 | |
| 63 | async def update_event_source( |
| 64 | event_source_id: int, |
| 65 | update_data: EventSourceUpdate, |
| 66 | db: AsyncSession, |
| 67 | ) -> EventSources: |
| 68 | event_source = await get_event_source_by_id(event_source_id, db) |
| 69 | event_source.update_from_model(update_data) |
| 70 | await db.commit() |
| 71 | await db.refresh(event_source) |
| 72 | return event_source |
| 73 | |
| 74 | |
| 75 | async def delete_event_source( |
| 76 | event_source_id: int, |
| 77 | db: AsyncSession, |
| 78 | ) -> None: |
| 79 | event_source = await get_event_source_by_id(event_source_id, db) |
| 80 | await db.delete(event_source) |
| 81 | await db.commit() |
| 82 | logger.info(f"Deleted event source {event_source_id}") |