"""Hourly source-location aggregation and timeline reads.""" from __future__ import annotations import asyncio import hashlib import hmac import logging import uuid from collections import defaultdict from datetime import datetime, timedelta, timezone from typing import Literal from sqlalchemy import delete, func, select from sqlalchemy.ext.asyncio import AsyncSession from app.core.config import settings from app.db.session import SessionLocal from app.models.permission_access_log import PermissionAccessLog from app.models.security_access_log import SecurityAccessLog from app.models.source_location_snapshot import SourceLocationSnapshot from app.services.geo_location_metadata import resolve_geo_location_metadata from app.services.ip_geolocation_fallback import resolve_external_ip_locations from app.services.ip_location import resolve_ip_location logger = logging.getLogger("ctms.source_location_aggregator") def _normalize_datetime(value: datetime) -> datetime: return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value.astimezone(timezone.utc) def _identity_hash(kind: str, value: str | None) -> str: if not value: return "" payload = f"source-location:{kind}:{value}".encode("utf-8") return hmac.new(settings.JWT_SECRET_KEY.encode("utf-8"), payload, hashlib.sha256).hexdigest() async def aggregate_source_location_hour(bucket_start: datetime, bucket_end: datetime) -> int: bucket_start = _normalize_datetime(bucket_start) bucket_end = _normalize_datetime(bucket_end) async with SessionLocal() as session: matching_security_request = select(SecurityAccessLog.id).where( PermissionAccessLog.request_id.is_not(None), SecurityAccessLog.request_id == PermissionAccessLog.request_id, ).exists() permission_rows = ( await session.execute( select( PermissionAccessLog.ip_address, PermissionAccessLog.user_id, func.count().filter(PermissionAccessLog.allowed.is_(True)).label("allowed_count"), func.count().filter(PermissionAccessLog.allowed.is_(False)).label("denied_count"), func.min(PermissionAccessLog.created_at).label("first_seen_at"), func.max(PermissionAccessLog.created_at).label("last_seen_at"), ) .where( PermissionAccessLog.created_at >= bucket_start, PermissionAccessLog.created_at < bucket_end, PermissionAccessLog.ip_address.is_not(None), ~matching_security_request, ) .group_by(PermissionAccessLog.ip_address, PermissionAccessLog.user_id) ) ).all() security_rows = ( await session.execute( select( SecurityAccessLog.client_ip, SecurityAccessLog.user_identifier, func.count().filter(SecurityAccessLog.status_code < 400).label("allowed_count"), func.count().filter(SecurityAccessLog.status_code >= 400).label("denied_count"), func.count() .filter( SecurityAccessLog.category.is_not(None), ~SecurityAccessLog.category.in_(("OTHER", "NOT_FOUND_NOISE")), ) .label("security_event_count"), func.count() .filter(SecurityAccessLog.severity.in_(("HIGH", "CRITICAL"))) .label("high_risk_count"), func.count() .filter( SecurityAccessLog.status_code >= 400, SecurityAccessLog.auth_status.in_(("INVALID_TOKEN", "ANONYMOUS")), ) .label("auth_failure_count"), func.min(SecurityAccessLog.created_at).label("first_seen_at"), func.max(SecurityAccessLog.created_at).label("last_seen_at"), ) .where( SecurityAccessLog.created_at >= bucket_start, SecurityAccessLog.created_at < bucket_end, SecurityAccessLog.client_ip.is_not(None), ) .group_by(SecurityAccessLog.client_ip, SecurityAccessLog.user_identifier) ) ).all() aggregates: dict[tuple[str, str], dict] = {} def merge_row( ip_address: str, user_identity: str, allowed_count: int, denied_count: int, first_seen_at: datetime, last_seen_at: datetime, *, security_event_count: int = 0, high_risk_count: int = 0, auth_failure_count: int = 0, ) -> None: key = (ip_address, user_identity) row = aggregates.setdefault( key, { "ip_address": ip_address, "user_identity": user_identity, "allowed_count": 0, "denied_count": 0, "security_event_count": 0, "high_risk_count": 0, "auth_failure_count": 0, "first_seen_at": _normalize_datetime(first_seen_at), "last_seen_at": _normalize_datetime(last_seen_at), }, ) row["allowed_count"] += int(allowed_count or 0) row["denied_count"] += int(denied_count or 0) row["security_event_count"] += int(security_event_count or 0) row["high_risk_count"] += int(high_risk_count or 0) row["auth_failure_count"] += int(auth_failure_count or 0) row["first_seen_at"] = min(row["first_seen_at"], _normalize_datetime(first_seen_at)) row["last_seen_at"] = max(row["last_seen_at"], _normalize_datetime(last_seen_at)) for ip_address, user_id, allowed, denied, first_seen, last_seen in permission_rows: merge_row(str(ip_address), str(user_id or ""), allowed, denied, first_seen, last_seen) for ip_address, user_identifier, allowed, denied, security_events, high_risk, auth_failures, first_seen, last_seen in security_rows: merge_row( str(ip_address), str(user_identifier or ""), allowed, denied, first_seen, last_seen, security_event_count=security_events, high_risk_count=high_risk, auth_failure_count=auth_failures, ) resolved_locations = { ip_address: resolve_ip_location(ip_address) for ip_address, _user_identity in aggregates } resolved_metadata = { ip_address: resolve_geo_location_metadata(ip_info) for ip_address, ip_info in resolved_locations.items() } missing_coordinate_ips = [ ip_address for ip_address, metadata in resolved_metadata.items() if metadata.accuracy_level != "private" and (metadata.longitude is None or metadata.latitude is None) ] external_locations = await resolve_external_ip_locations(missing_coordinate_ips) for ip_address, external_location in external_locations.items(): resolved_locations[ip_address] = external_location.merge_ip_location(resolved_locations[ip_address]) resolved_metadata[ip_address] = external_location.to_metadata() await session.execute( delete(SourceLocationSnapshot).where(SourceLocationSnapshot.bucket_time == bucket_start) ) for row in aggregates.values(): ip_info = resolved_locations[row["ip_address"]] metadata = resolved_metadata[row["ip_address"]] longitude = metadata.longitude latitude = metadata.latitude country = metadata.country or ip_info.country location = " / ".join( part for part in [country, ip_info.province, ip_info.city] if part ) or ip_info.location or "未知" session.add( SourceLocationSnapshot( id=uuid.uuid4(), bucket_time=bucket_start, ip_hash=_identity_hash("ip", row["ip_address"]), user_hash=_identity_hash("user", row["user_identity"]), country=country, country_code=metadata.country_code, province=ip_info.province, region_code=metadata.region_code, city=ip_info.city, isp=ip_info.isp, location=location, longitude=longitude, latitude=latitude, accuracy_level=metadata.accuracy_level, allowed_count=row["allowed_count"], denied_count=row["denied_count"], security_event_count=row["security_event_count"], high_risk_count=row["high_risk_count"], auth_failure_count=row["auth_failure_count"], first_seen_at=row["first_seen_at"], last_seen_at=row["last_seen_at"], ) ) await session.commit() return len(aggregates) async def get_source_location_timeline( db: AsyncSession, *, start_at: datetime, end_at: datetime, granularity: Literal["hour", "day"], ) -> list[dict]: rows = ( await db.execute( select(SourceLocationSnapshot) .where( SourceLocationSnapshot.bucket_time >= start_at, SourceLocationSnapshot.bucket_time < end_at, ) .order_by(SourceLocationSnapshot.bucket_time) ) ).scalars().all() buckets: dict[datetime, dict] = defaultdict( lambda: { "allowed_count": 0, "denied_count": 0, "security_event_count": 0, "high_risk_count": 0, "ip_hashes": set(), "user_hashes": set(), } ) for row in rows: bucket = _normalize_datetime(row.bucket_time) if granularity == "day": bucket = bucket.replace(hour=0, minute=0, second=0, microsecond=0) else: bucket = bucket.replace(minute=0, second=0, microsecond=0) item = buckets[bucket] item["allowed_count"] += row.allowed_count item["denied_count"] += row.denied_count item["security_event_count"] += row.security_event_count item["high_risk_count"] += row.high_risk_count item["ip_hashes"].add(row.ip_hash) if row.user_hash: item["user_hashes"].add(row.user_hash) return [ { "bucket_time": bucket.isoformat(), "total_count": item["allowed_count"] + item["denied_count"], "allowed_count": item["allowed_count"], "denied_count": item["denied_count"], "security_event_count": item["security_event_count"], "high_risk_count": item["high_risk_count"], "unique_ip_count": len(item["ip_hashes"]), "unique_user_count": len(item["user_hashes"]), } for bucket, item in sorted(buckets.items()) ] async def run_hourly_source_location_aggregation(stop_event: asyncio.Event) -> None: logger.info("Source location aggregator started") now = datetime.now(timezone.utc) completed_hour = now.replace(minute=0, second=0, microsecond=0) try: await aggregate_source_location_hour(completed_hour - timedelta(hours=1), completed_hour) except Exception: logger.exception("Failed to backfill source location snapshot for %s", completed_hour) while not stop_event.is_set(): now = datetime.now(timezone.utc) next_hour = now.replace(minute=0, second=0, microsecond=0) + timedelta(hours=1) try: await asyncio.wait_for(stop_event.wait(), timeout=(next_hour - now).total_seconds()) break except asyncio.TimeoutError: pass try: await aggregate_source_location_hour(next_hour - timedelta(hours=1), next_hour) except Exception: logger.exception("Failed to aggregate source locations for %s", next_hour) logger.info("Source location aggregator stopped")