Skip to content
Merged
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
34 changes: 32 additions & 2 deletions backend/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,25 @@ def short_ts(ts):
templates.env.filters["short_ts"] = short_ts


def normalize_asn(asn: str | None) -> str | None:
"""Return a normalized ASN string starting with ``AS``.

Various geo-IP providers return the autonomous system number in different
formats (e.g. ``"AS906"``, ``"906"`` or ``"AS906 Network"``). For the
purpose of de-duplicating records we treat ``AS906`` and ``906`` as the same
ASN. This helper extracts the numeric portion and ensures the value is
consistently prefixed with ``AS``.
"""

if not asn:
return asn

match = re.search(r"(\d+)", str(asn))
if match:
return f"AS{match.group(1)}"
return str(asn).upper()


@app.on_event("startup")
def create_default_user():
db = database.SessionLocal()
Expand Down Expand Up @@ -388,6 +407,7 @@ def probe_page(request: Request, db: Session = Depends(get_db)):
except Exception:
pass

data["asn"] = normalize_asn(data.get("asn"))
db_record = models.TestRecord(**data)
db.add(db_record)
db.commit()
Expand Down Expand Up @@ -456,6 +476,9 @@ def read_tests(request: Request, db: Session = Depends(get_db)):
.all()
)

for r in rows:
r.asn = normalize_asn(r.asn)

if not rows:
return {"message": "No recent test records found", "records": []}

Expand Down Expand Up @@ -523,6 +546,7 @@ def create_test(
except Exception:
pass

data["asn"] = normalize_asn(data.get("asn"))
if not skip_ping:
if not data.get("ping_ms"):
host = data.get("test_target") or client_ip
Expand Down Expand Up @@ -585,7 +609,7 @@ def create_test(
averaged = {
"client_ip": client_ip,
"location": data.get("location") or existing_records[0].location,
"asn": data.get("asn") or existing_records[0].asn,
"asn": normalize_asn(data.get("asn") or existing_records[0].asn),
"isp": data.get("isp") or existing_records[0].isp,
"ping_min_ms": sum(values_ping_min) / len(values_ping_min)
if values_ping_min
Expand Down Expand Up @@ -621,6 +645,8 @@ def admin_read_tests(
db: Session = Depends(get_db), user: models.User = Depends(require_active_user)
):
records = db.query(models.TestRecord).all()
for r in records:
r.asn = normalize_asn(r.asn)
return {"records": records}


Expand All @@ -632,7 +658,9 @@ def admin_create_test(
db: Session = Depends(get_db),
user: models.User = Depends(require_active_user),
):
db_record = models.TestRecord(**record.dict())
data = record.dict()
data["asn"] = normalize_asn(data.get("asn"))
db_record = models.TestRecord(**data)
db.add(db_record)
db.commit()
db.refresh(db_record)
Expand All @@ -652,6 +680,8 @@ def admin_update_test(
if not db_record:
raise HTTPException(status_code=404, detail="Record not found")
for key, value in record.dict(exclude_unset=True).items():
if key == "asn":
value = normalize_asn(value)
setattr(db_record, key, value)
db.commit()
db.refresh(db_record)
Expand Down
44 changes: 44 additions & 0 deletions backend/test_main.py
Original file line number Diff line number Diff line change
Expand Up @@ -261,3 +261,47 @@ def test_multi_speedtest_preserves_single_record():
assert speeds["single"]["download_mbps"] == 10
assert abs(speeds["multi"]["download_mbps"] - 50) < 0.01


def test_asn_normalization_and_merge():
from backend.database import SessionLocal
from backend.models import TestRecord

db = SessionLocal()
try:
db.query(TestRecord).delete()
db.commit()
finally:
db.close()

headers = {"X-Forwarded-For": "8.8.8.8"}
payload1 = {
"asn": "906",
"isp": "ISP1",
"location": "Loc1",
"ping_ms": 10,
"ping_min_ms": 10,
"ping_max_ms": 10,
}
client.post("/tests?skip_ping=true", json=payload1, headers=headers)

payload2 = {
"asn": "AS906",
"isp": "ISP2",
"location": "Loc2",
"ping_ms": 20,
"ping_min_ms": 20,
"ping_max_ms": 20,
}
client.post("/tests?skip_ping=true", json=payload2, headers=headers)

res = client.get("/tests")
assert res.status_code == 200
data = res.json()
records = [r for r in data["records"] if r["client_ip"] == "8.8.8.8"]
assert len(records) == 1
rec = records[0]
assert rec["asn"] == "AS906"
assert abs(rec["ping_ms"] - 15) < 0.01
assert abs(rec["ping_min_ms"] - 15) < 0.01
assert abs(rec["ping_max_ms"] - 15) < 0.01