app/probes/__init__.py (view raw)
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 |
from __future__ import annotations
import asyncio
import logging
import random
from datetime import UTC, datetime, timedelta
import httpx
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import Settings
from app.db.models import DnsRecord, ProbeResult, Tracker
from app.probes.http import probe_http
from app.probes.udp import probe_udp
from app.services.dns import host_has_ipv6, lookup_asn, resolve_cname_chain
logger = logging.getLogger(__name__)
def compute_next_check(
*,
settings: Settings,
announced_interval: int | None,
consecutive_failures: int,
success: bool,
) -> datetime:
now = datetime.now(UTC)
if success:
base = announced_interval or settings.default_probe_interval_seconds
base = max(base, settings.min_probe_interval_seconds)
else:
exp = min(consecutive_failures, 12)
base = min(
settings.min_probe_interval_seconds * (2 ** max(exp - 1, 0)),
12 * 3600,
)
base = max(base, settings.min_probe_interval_seconds)
jitter = random.uniform(0, min(60.0, base * 0.1))
return now + timedelta(seconds=base + jitter)
async def _upsert_dns_records(
session: AsyncSession, tracker: Tracker, records: list[tuple[str, str]]
) -> None:
now = datetime.now(UTC)
for rtype, value in records:
result = await session.execute(
select(DnsRecord).where(
DnsRecord.tracker_id == tracker.id,
DnsRecord.record_type == rtype,
DnsRecord.value == value,
)
)
row = result.scalar_one_or_none()
if row is None:
session.add(
DnsRecord(
tracker_id=tracker.id,
record_type=rtype,
value=value,
first_seen_at=now,
last_seen_at=now,
)
)
else:
row.last_seen_at = now
async def probe_tracker(
session: AsyncSession,
tracker: Tracker,
settings: Settings,
client: httpx.AsyncClient,
) -> ProbeResult:
now = datetime.now(UTC)
dns = await resolve_cname_chain(tracker.hostname)
tracker.terminal_cname = dns.terminal_cname
tracker.infrastructure_fingerprint = dns.fingerprint
dns_rows: list[tuple[str, str]] = [("CNAME", c) for c in dns.cname_chain]
dns_rows.extend(("A", ip) for ip in dns.ipv4)
dns_rows.extend(("AAAA", ip) for ip in dns.ipv6)
await _upsert_dns_records(session, tracker, dns_rows)
ipv6_mode = settings.ipv6_mode()
can_v6 = host_has_ipv6() if ipv6_mode == "auto" else ipv6_mode == "true"
tracker.supports_ipv6 = "not_tested" if not can_v6 else ("true" if dns.ipv6 else "false")
tracker.supports_ipv4 = bool(dns.ipv4) if settings.enable_ipv4 else None
public_ips = list(dns.ipv4)
if can_v6:
public_ips.extend(dns.ipv6)
if settings.geoip_asn_db and public_ips:
asn, net, country = lookup_asn(public_ips[0], settings.geoip_asn_db)
tracker.asn = asn
tracker.network_name = net
tracker.country_code = country
if tracker.scheme == "udp":
peer = (settings.peer_id_prefix() + b"xxxxxxxxxxxx")[:20]
outcome = await probe_udp(
tracker.hostname,
tracker.port,
timeout=settings.probe_timeout_seconds,
peer_id=peer,
)
elif tracker.scheme in {"http", "https"}:
outcome = await probe_http(
tracker.canonical_url,
client=client,
timeout=settings.probe_timeout_seconds,
max_bytes=settings.max_response_bytes,
peer_id_prefix=settings.peer_id_prefix(),
resolved_ips=public_ips or None,
)
else:
tracker.current_status = "unsupported"
tracker.last_checked_at = now
tracker.next_check_at = now + timedelta(days=30)
result = ProbeResult(
tracker_id=tracker.id,
checked_at=now,
status="unsupported",
response_valid=False,
error_kind="unsupported",
error_detail="scheme not probed",
)
session.add(result)
await session.commit()
return result
success = outcome.response_valid and outcome.status == "up"
if success:
tracker.consecutive_failures = 0
tracker.current_status = "up"
if outcome.tracker_interval_seconds:
tracker.announced_interval_seconds = outcome.tracker_interval_seconds
else:
tracker.consecutive_failures += 1
tracker.current_status = (
outcome.status if outcome.status in {"down", "degraded"} else "down"
)
tracker.last_latency_ms = outcome.latency_ms
tracker.last_checked_at = now
tracker.next_check_at = compute_next_check(
settings=settings,
announced_interval=tracker.announced_interval_seconds,
consecutive_failures=tracker.consecutive_failures,
success=success,
)
result = ProbeResult(
tracker_id=tracker.id,
checked_at=now,
status=tracker.current_status,
latency_ms=outcome.latency_ms,
ip_address=outcome.ip_address,
ip_family=outcome.ip_family,
response_valid=outcome.response_valid,
tracker_interval_seconds=outcome.tracker_interval_seconds,
seeders=outcome.seeders,
leechers=outcome.leechers,
error_kind=outcome.error_kind,
error_detail=(outcome.error_detail or "")[:512] or None,
)
session.add(result)
await session.commit()
return result
async def probe_due_batch(
session: AsyncSession,
settings: Settings,
*,
limit: int | None = None,
) -> int:
from app.db import get_session_factory
from app.db.queries import due_trackers
limit = limit or settings.probe_concurrency
trackers = await due_trackers(session, limit=limit)
if not trackers:
return 0
ids = [t.id for t in trackers]
sem = asyncio.Semaphore(settings.probe_concurrency)
factory = get_session_factory()
async with httpx.AsyncClient(
timeout=settings.probe_timeout_seconds,
headers={"User-Agent": settings.user_agent()},
follow_redirects=False,
) as client:
async def _one(tracker_id: int) -> None:
async with sem, factory() as own_session:
tracker = await own_session.get(Tracker, tracker_id)
if tracker is None:
return
try:
await probe_tracker(own_session, tracker, settings, client)
except Exception: # noqa: BLE001
logger.exception("probe failed for id=%s", tracker_id)
await asyncio.gather(*[_one(i) for i in ids])
return len(ids)
|