Files
ctms/backend/app/services/source_location_aggregator.py
T
Cheng Zhou d5279b124f
Storage Persistence Guard / storage-persistence-audit (push) Has been cancelled
Client Quality Gates / Shared client and Web (push) Has been cancelled
Client Quality Gates / macOS Desktop (push) Has been cancelled
Client Quality Gates / Shared client and Web (pull_request) Has been cancelled
Client Quality Gates / macOS Desktop (pull_request) Has been cancelled
Storage Persistence Guard / storage-persistence-audit (pull_request) Has been cancelled
release(main): 同步 dev 最新候选改动
2026-07-16 17:15:50 +08:00

293 lines
12 KiB
Python

"""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")