diff options
Diffstat (limited to 'packages')
6 files changed, 484 insertions, 84 deletions
diff --git a/packages/meshbay-hub/src/meshbay_hub/tasks/cleanup.py b/packages/meshbay-hub/src/meshbay_hub/tasks/cleanup.py index b0ef4d4..559d872 100644 --- a/packages/meshbay-hub/src/meshbay_hub/tasks/cleanup.py +++ b/packages/meshbay-hub/src/meshbay_hub/tasks/cleanup.py @@ -4,7 +4,7 @@ import asyncio import logging from datetime import UTC, datetime, timedelta -from sqlalchemy import delete, select +from sqlalchemy import delete, select, update from sqlalchemy.ext.asyncio import AsyncSession from meshbay_hub.db.models import EmailVerification, Group, GroupInviteLink, IPLog, User @@ -35,9 +35,49 @@ async def purge_expired_verifications(db: AsyncSession) -> int: async def purge_stale_pending_users(db: AsyncSession, expiry_days: int = PENDING_USER_EXPIRY_DAYS) -> int: + """Delete accounts that never verified their address, and what points at them. + + A pending account is not childless: registering writes an `account_create` + IP-log row, a failed sign-in another, and a group owner may have added it to + a group. A bare `DELETE FROM users` therefore violates those foreign keys — + which PostgreSQL enforces and the SQLite the suite runs on does not — so on + the real hub it raised, the stale account stayed, and every later run + raised again. + + The IP log is kept and detached, keeping the name, exactly as + `erase_account` does for a deleted account: it is the legal record. Every + other nullable reference is cleared and every other row deleted, found from + the schema so a table added later is covered. An account that owns a group + is left alone — a pending account cannot create one, so that would be a + state this function did not expect and should not guess about. + """ + from meshbay_hub.db.models import Base, Group + cutoff = datetime.now(UTC) - timedelta(days=expiry_days) - result = await db.execute( - delete(User).where(User.status == "pending", User.created_at < cutoff)) + stale = (await db.execute( + select(User.id, User.username).where( + User.status == "pending", User.created_at < cutoff, + ~User.id.in_(select(Group.admin_id))))).all() + if not stale: + return 0 + ids = [uid for uid, _ in stale] + + for uid, name in stale: + await db.execute(update(IPLog).where(IPLog.user_id == uid) + .values(username=name, user_id=None)) + for table in Base.metadata.sorted_tables: + if table.name in (User.__tablename__, IPLog.__tablename__): + continue + for fk in table.foreign_keys: + if fk.column.table.name != User.__tablename__: + continue + column = fk.parent + if column.nullable: + await db.execute(update(table).where(column.in_(ids)) + .values({column.name: None})) + else: + await db.execute(delete(table).where(column.in_(ids))) + result = await db.execute(delete(User).where(User.id.in_(ids))) await db.commit() return result.rowcount @@ -53,42 +93,57 @@ async def purge_invite_links(db: AsyncSession) -> int: return result.rowcount +async def _purge_login_throttle(db: AsyncSession) -> int: + from meshbay_hub import login_throttle + return await login_throttle.purge_expired(db) + + +async def _purge_mail_quota(db: AsyncSession) -> int: + from meshbay_hub import mail + return await mail.purge_expired_quota(db) + + +# In order, each in its own session and its own try: one step that raises must +# not cost the others their run. They used to share one `try`, so the day +# `purge_stale_pending_users` first hit a foreign key on PostgreSQL, the mail +# counters, the sign-in counters and the invitation links behind it stopped +# being purged at all — silently, since the loop logged one line and slept. +_STEPS = ( + ("IP log entries older than the retention period", purge_old_ip_logs), + ("expired email verifications", purge_expired_verifications), + ("stale pending users", purge_stale_pending_users), + # One row per recipient the hub has written to, and the window is a day: + # without this the table grows for the life of the instance. + ("expired mail counters", _purge_mail_quota), + # Every name anybody types at the sign-in form is a row, real or not. + ("expired sign-in counters", _purge_login_throttle), + ("spent or expired invitation links", purge_invite_links), +) + + +async def run_cleanup(get_session) -> dict[str, int | None]: + """One pass of every step. A step that failed reports None.""" + done: dict[str, int | None] = {} + for what, step in _STEPS: + try: + async with get_session() as db: + n = await step(db) + done[what] = n + if n: + log.info("Purged %d %s", n, what) + except asyncio.CancelledError: + raise + except Exception as e: + done[what] = None + log.error("Cleanup of %s failed: %s", what, e) + return done + + async def cleanup_loop(get_session): """Run cleanup once at startup, then every 24 hours.""" try: while True: - try: - async with get_session() as db: - deleted = await purge_old_ip_logs(db) - if deleted: - log.info("Purged %d IP log entries older than %d days", - deleted, RETENTION_DAYS) - expired = await purge_expired_verifications(db) - if expired: - log.info("Purged %d expired email verifications", expired) - stale = await purge_stale_pending_users(db) - if stale: - log.info("Purged %d stale pending users", stale) - # One row per recipient the hub has written to, and the - # window is a day: without this the table grows for the - # life of the instance and nothing ever reads the old rows. - from meshbay_hub import mail - quota = await mail.purge_expired_quota(db) - if quota: - log.info("Purged %d expired mail counters", quota) - # Every name anybody types at the sign-in form is a row, - # real or not; once its window has passed, nothing reads it. - from meshbay_hub import login_throttle - throttled = await login_throttle.purge_expired(db) - if throttled: - log.info("Purged %d expired sign-in counters", throttled) - links = await purge_invite_links(db) - if links: - log.info("Purged %d spent or expired invitation links", links) - except asyncio.CancelledError: - raise - except Exception as e: - log.error("IP log cleanup failed: %s", e) + await run_cleanup(get_session) await asyncio.sleep(CLEANUP_INTERVAL_HOURS * 3600) except asyncio.CancelledError: return diff --git a/packages/meshbay-hub/tests/test_cleanup_foreign_keys.py b/packages/meshbay-hub/tests/test_cleanup_foreign_keys.py new file mode 100644 index 0000000..b4994b2 --- /dev/null +++ b/packages/meshbay-hub/tests/test_cleanup_foreign_keys.py @@ -0,0 +1,103 @@ +""" +The daily cleanup, against a database that enforces foreign keys. + +PostgreSQL always does; the SQLite the rest of the suite runs on does not unless +asked. So `DELETE FROM users` for an account that still had an IP-log row passed +every test here and raised on the real hub — and because every step shared one +`try`, the purges behind it (mail counters, sign-in counters, invitation links) +stopped running for good. Both halves are held here, on an engine with +`PRAGMA foreign_keys=ON`. +""" + +from datetime import UTC, datetime, timedelta + +import pytest +from meshbay_hub.db.models import ( + Base, + EmailVerification, + Group, + GroupMember, + IPLog, + Notification, + User, +) +from meshbay_hub.tasks import cleanup +from sqlalchemy import event, select +from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine + + +@pytest.fixture +async def strict_db(tmp_path): + engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'hub.db'}") + + @event.listens_for(engine.sync_engine, "connect") + def _enforce(dbapi_conn, _record): + dbapi_conn.execute("PRAGMA foreign_keys=ON") + + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + yield async_sessionmaker(engine, expire_on_commit=False) + await engine.dispose() + + +def _user(name, status, age_days): + return User(username=name, email="x", pw_hash=b"x", pw_salt=b"x", hub_id="h", + status=status, + created_at=datetime.now(UTC) - timedelta(days=age_days)) + + +async def test_a_stale_pending_account_is_purged_with_what_points_at_it(strict_db): + async with strict_db() as db: + owner = _user("groupowner", "active", 30) + stale = _user("neververified", "pending", 8) + db.add_all([owner, stale]) + await db.flush() + group = Group(name="g", admin_id=owner.id) + db.add(group) + await db.flush() + db.add_all([ + IPLog(user_id=stale.id, event="account_create", ip_address="192.0.2.1"), + GroupMember(group_id=group.id, user_id=stale.id), + Notification(user_id=stale.id, kind="group_invite", title="t"), + EmailVerification(email_hash="h", code="1", purpose="registration", + user_id=stale.id, + expires_at=datetime.now(UTC) - timedelta(days=6)), + ]) + await db.commit() + stale_id = stale.id + + async with strict_db() as db: + assert await cleanup.purge_stale_pending_users(db) == 1 + + async with strict_db() as db: + assert await db.get(User, stale_id) is None + log_row = (await db.execute(select(IPLog))).scalar_one() + # The legal record survives, still saying who it was about. + assert log_row.user_id is None and log_row.username == "neververified" + assert (await db.execute(select(GroupMember).where( + GroupMember.user_id == stale_id))).first() is None + + +async def test_a_recent_or_active_account_is_left_alone(strict_db): + async with strict_db() as db: + db.add_all([_user("stillpending", "pending", 2), + _user("activeuser", "active", 30)]) + await db.commit() + async with strict_db() as db: + assert await cleanup.purge_stale_pending_users(db) == 0 + + +async def test_one_failing_step_does_not_stop_the_others(strict_db, monkeypatch): + ran = [] + + async def broken(db): + raise RuntimeError("boom") + + async def later(db): + ran.append("later") + return 0 + + monkeypatch.setattr(cleanup, "_STEPS", (("broken", broken), ("later", later))) + done = await cleanup.run_cleanup(strict_db) + assert done == {"broken": None, "later": 0} + assert ran == ["later"] diff --git a/packages/meshbay-node/src/meshbay_node/linkpreview.py b/packages/meshbay-node/src/meshbay_node/linkpreview.py index b223aea..6e1e618 100644 --- a/packages/meshbay-node/src/meshbay_node/linkpreview.py +++ b/packages/meshbay-node/src/meshbay_node/linkpreview.py @@ -30,6 +30,7 @@ OG image rides the existing `media_cache` thumb store (same as a poster). from __future__ import annotations +import asyncio import ipaddress import logging import socket @@ -42,6 +43,9 @@ import httpx log = logging.getLogger(__name__) _TIMEOUT = 5.0 +# The whole fetch, redirects and body included. `_TIMEOUT` is per network +# operation, so a server that keeps sending slowly never trips it on its own. +_TOTAL_DEADLINE = 15.0 _MAX_REDIRECTS = 3 _MAX_HTML_BYTES = 512 * 1024 _MAX_IMAGE_BYTES = 2 * 1024 * 1024 @@ -171,19 +175,56 @@ def _reject_if_rebound(resp: httpx.Response) -> None: async def _get(client: httpx.AsyncClient, url: str) -> httpx.Response: - """One GET with manual, re-validated redirects.""" + """ + One GET with manual, re-validated redirects, **body not read**. + + The caller reads it through `_read_capped` and must close it. `client.get` + is not usable here: it reads and decodes the whole body before returning, + so the size caps applied afterwards bounded nothing — a page, an image or a + compressed stream of any size was held in memory first. The caps are the + only thing between a URL a member pasted and the node's memory. + """ current = safe_url(url) for _ in range(_MAX_REDIRECTS + 1): - resp = await client.get(current, headers={"User-Agent": _UA}, - follow_redirects=False) - _reject_if_rebound(resp) + request = client.build_request("GET", current, headers={"User-Agent": _UA}) + resp = await client.send(request, stream=True, follow_redirects=False) + try: + _reject_if_rebound(resp) + except Exception: + await resp.aclose() + raise if resp.is_redirect and "location" in resp.headers: - current = safe_url(urljoin(current, resp.headers["location"])) + location = resp.headers["location"] + await resp.aclose() + current = safe_url(urljoin(current, location)) continue return resp raise UnsafeURL("too many redirects") +async def _read_capped(resp: httpx.Response, cap: int) -> tuple[bytes, bool]: + """ + At most `cap` bytes of the *decoded* body, and whether there was more. + + Counted after decoding, so a small compressed response that inflates to + gigabytes stops at the cap like any other. A declared length over the cap + is not read at all. + """ + try: + declared = int(resp.headers.get("content-length", "")) + except ValueError: + declared = -1 + if resp.headers.get("content-encoding", "identity") == "identity" \ + and declared > cap: + return b"", True + body = bytearray() + async for chunk in resp.aiter_bytes(): + body += chunk + if len(body) > cap: + return bytes(body[:cap]), True + return bytes(body), False + + async def fetch_preview(url: str, *, client: httpx.AsyncClient | None = None) -> dict | None: """ Return {url, title, description, site_name, image_url} for a URL, or None @@ -195,17 +236,25 @@ async def fetch_preview(url: str, *, client: httpx.AsyncClient | None = None) -> if own: client = httpx.AsyncClient(timeout=_TIMEOUT, max_redirects=0) try: - safe_url(url) - resp = await _get(client, url) + return await asyncio.wait_for(_preview(client, url), _TOTAL_DEADLINE) + except (httpx.HTTPError, UnsafeURL, TimeoutError) as e: + log.debug("link preview for %s: %s", url[:80], e) + return None + finally: + if own: + await client.aclose() + + +async def _preview(client: httpx.AsyncClient, url: str) -> dict | None: + safe_url(url) + resp = await _get(client, url) + try: ctype = resp.headers.get("content-type", "").split(";")[0].strip().lower() if resp.status_code != 200 or ctype not in ("text/html", "application/xhtml+xml"): return None - - body = b"" - async for chunk in resp.aiter_bytes(): - body += chunk - if len(body) >= _MAX_HTML_BYTES: - break + # A page longer than this is cut, not refused: what a card needs is in + # the head. + body, _truncated = await _read_capped(resp, _MAX_HTML_BYTES) final_url = str(resp.url) parser = _HeadParser() @@ -237,12 +286,8 @@ async def fetch_preview(url: str, *, client: httpx.AsyncClient | None = None) -> "site_name": (site_name or "")[:120] or None, "image_url": image, } - except (httpx.HTTPError, UnsafeURL) as e: - log.debug("link preview for %s: %s", url[:80], e) - return None finally: - if own: - await client.aclose() + await resp.aclose() async def fetch_image(url: str, *, client: httpx.AsyncClient | None = None) -> bytes | None: @@ -251,23 +296,30 @@ async def fetch_image(url: str, *, client: httpx.AsyncClient | None = None) -> b if own: client = httpx.AsyncClient(timeout=_TIMEOUT, max_redirects=0) try: - safe_url(url) - resp = await _get(client, url) - ctype = resp.headers.get("content-type", "").split(";")[0].strip().lower() - if resp.status_code != 200 or not ctype.startswith("image/"): - return None - raw = b"" - async for chunk in resp.aiter_bytes(): - raw += chunk - if len(raw) > _MAX_IMAGE_BYTES: - return None - return _downscale(raw) - except (httpx.HTTPError, UnsafeURL) as e: + raw = await asyncio.wait_for(_image_bytes(client, url), _TOTAL_DEADLINE) + except (httpx.HTTPError, UnsafeURL, TimeoutError) as e: log.debug("link preview image %s: %s", url[:80], e) return None finally: if own: await client.aclose() + if raw is None: + return None + # Decoding an image is CPU work on untrusted bytes; not on the event loop. + return await asyncio.to_thread(_downscale, raw) + + +async def _image_bytes(client: httpx.AsyncClient, url: str) -> bytes | None: + safe_url(url) + resp = await _get(client, url) + try: + ctype = resp.headers.get("content-type", "").split(";")[0].strip().lower() + if resp.status_code != 200 or not ctype.startswith("image/"): + return None + raw, too_big = await _read_capped(resp, _MAX_IMAGE_BYTES) + return None if too_big else raw + finally: + await resp.aclose() def _downscale(raw: bytes) -> bytes | None: diff --git a/packages/meshbay-node/src/meshbay_node/transport/webrtc/admission.py b/packages/meshbay-node/src/meshbay_node/transport/webrtc/admission.py index e678a13..f94cade 100644 --- a/packages/meshbay-node/src/meshbay_node/transport/webrtc/admission.py +++ b/packages/meshbay-node/src/meshbay_node/transport/webrtc/admission.py @@ -40,9 +40,18 @@ _INVITE_ID_RE = re.compile(r"[0-9a-f]{32}") MAX_JOIN_ATTEMPTS = 5 # Per-connection limits alone would not bind an attacker who can open connections # at will — and the adversary who can mint tokens for any account is the hub. So -# failed pairings are also counted node-wide over a window. +# wrong codes are also counted per account and node-wide over a window. +# +# Only a *wrong code* counts, and the lock is consulted only when a code is about +# to be tried. Every member reconnecting obtains the group key through this very +# message — no per-member bundle is stored — so a lock checked at the top of it, +# fed by ordinary refusals such as `code_required`, let any one member refuse the +# key to everybody on the node for as long as they kept failing. A device the +# node already pinned, joining without a code, is never subject to it. MAX_JOIN_FAILURES_WINDOW = 20 +MAX_JOIN_FAILURES_PER_ACCOUNT = 5 JOIN_FAILURE_WINDOW = 600 # seconds +_JOIN_FAILURE_ACCOUNTS_TRACKED = 1000 class AdmissionMixin: @@ -124,15 +133,11 @@ class AdmissionMixin: # ── Pairing and join (H3, M3) ──────────────────────────────────────────── - def _join_refuse(self, reason: str, audit_detail: str = "") -> None: + def _join_refuse(self, reason: str, audit_detail: str = "", *, + wrong_code: bool = False) -> None: self._join_attempts += 1 - # Node-wide window, shared across connections: reconnecting must not reset - # the budget. - now = time.time() - failures = [t for t in self._ctx.get("join_failures", []) - if now - t < JOIN_FAILURE_WINDOW] - failures.append(now) - self._ctx["join_failures"] = failures + if wrong_code: + self._record_wrong_code() self._audit_join("join_refused", audit_detail or reason) self._send({ "type": MNP.JOIN_RESULT, @@ -141,6 +146,41 @@ class AdmissionMixin: "reason": reason, }) + def _record_wrong_code(self) -> None: + """Count one wrong code, per account and node-wide, across connections: + reconnecting must not reset either budget.""" + now = time.time() + failures = [t for t in self._ctx.get("join_failures", []) + if now - t < JOIN_FAILURE_WINDOW] + failures.append(now) + self._ctx["join_failures"] = failures + + per_account: dict = self._ctx.setdefault("join_failures_by_account", {}) + if len(per_account) > _JOIN_FAILURE_ACCOUNTS_TRACKED: + for uid, times in list(per_account.items()): + if not times or now - times[-1] >= JOIN_FAILURE_WINDOW: + per_account.pop(uid, None) + uid = self._user_id or getattr(self, "_pending_sub", "") + mine = [t for t in per_account.get(uid, ()) if now - t < JOIN_FAILURE_WINDOW] + mine.append(now) + per_account[uid] = mine + + def _code_guessing_locked(self, user_id: str) -> bool: + """Whether a code may be tried now. Refuses, and says so, when not.""" + now = time.time() + recent = [t for t in self._ctx.get("join_failures", []) + if now - t < JOIN_FAILURE_WINDOW] + mine = [t for t in (self._ctx.get("join_failures_by_account") or {}) + .get(user_id, ()) if now - t < JOIN_FAILURE_WINDOW] + if len(mine) >= MAX_JOIN_FAILURES_PER_ACCOUNT: + self._audit_join("join_throttled", f"{len(mine)} wrong codes by this account") + elif len(recent) >= MAX_JOIN_FAILURES_WINDOW: + self._audit_join("join_throttled", f"{len(recent)} wrong codes on this node") + else: + return False + self._send({"type": "error", "detail": "Pairing temporarily locked"}) + return True + def _audit_join(self, event: str, detail: str) -> None: audit = self._ctx.get("audit_store") if not audit: @@ -177,14 +217,6 @@ class AdmissionMixin: self._send({"type": "error", "detail": "Too many attempts"}) return - now = time.time() - recent = [t for t in self._ctx.get("join_failures", []) - if now - t < JOIN_FAILURE_WINDOW] - if len(recent) >= MAX_JOIN_FAILURES_WINDOW: - self._audit_join("join_throttled", f"{len(recent)} failures in window") - self._send({"type": "error", "detail": "Pairing temporarily locked"}) - return - user_id = self._user_id or getattr(self, "_pending_sub", "") username = self._username or getattr(self, "_pending_username", "") if not user_id: @@ -295,9 +327,11 @@ class AdmissionMixin: if not code: self._join_refuse("code_required") return + if self._code_guessing_locked(user_id): + return invite = await roster.consume_invite(code, user_id, session_group) if not invite: - self._join_refuse("code_invalid") + self._join_refuse("code_invalid", wrong_code=True) return await roster.set_member( group_id=invite["group_id"], user_id=user_id, @@ -337,9 +371,11 @@ class AdmissionMixin: self._join_refuse("code_required") return + if self._code_guessing_locked(user_id): + return invite = await roster.consume_invite(code, user_id, session_group) if not invite: - self._join_refuse("code_invalid") + self._join_refuse("code_invalid", wrong_code=True) return await self._pin_and_admit( diff --git a/packages/meshbay-node/tests/test_linkpreview.py b/packages/meshbay-node/tests/test_linkpreview.py index a14b173..9fca186 100644 --- a/packages/meshbay-node/tests/test_linkpreview.py +++ b/packages/meshbay-node/tests/test_linkpreview.py @@ -193,3 +193,83 @@ async def test_fetch_image_refuses_a_decompression_bomb(resolves_public, monkeyp return httpx.Response(200, headers={"content-type": "image/png"}, content=bomb) async with _client(handler) as c: assert await linkpreview.fetch_image("https://example.com/x.png", client=c) is None + + +# ── Size caps bind while reading, not after ───────────────────────────────── +# +# `client.get` read and decoded the whole body before the caps looked at it, so +# a link posted in chat could make the node hold any amount of data. These +# count what the server actually had to hand over. + +_CHUNK = 64 * 1024 + + +def _endless(counter, head=b""): + async def body(): + if head: + counter["sent"] += len(head) + yield head + for _ in range(2000): # 125 MiB if nobody stops + counter["sent"] += _CHUNK + yield b"x" * _CHUNK + return body() + + +async def test_a_huge_page_is_not_read_past_the_cap(resolves_public): + counter = {"sent": 0} + head = b"<html><head><title>T</title></head><body>" + + def handler(request): + return httpx.Response(200, headers={"content-type": "text/html"}, + content=_endless(counter, head)) + async with _client(handler) as c: + meta = await linkpreview.fetch_preview("https://example.com/", client=c) + assert meta is not None and meta["title"] == "T" + assert counter["sent"] <= linkpreview._MAX_HTML_BYTES + 2 * _CHUNK + + +async def test_a_huge_image_is_refused_without_being_read(resolves_public): + counter = {"sent": 0} + + def handler(request): + return httpx.Response(200, headers={"content-type": "image/png"}, + content=_endless(counter)) + async with _client(handler) as c: + assert await linkpreview.fetch_image("https://example.com/x.png", client=c) is None + assert counter["sent"] <= linkpreview._MAX_IMAGE_BYTES + 2 * _CHUNK + + +async def test_a_compressed_body_is_capped_after_decoding(resolves_public): + import zlib + comp = zlib.compressobj(9, zlib.DEFLATED, 31) # gzip + counter = {"inflated": 0} + + async def body(): + yield comp.compress(b"<html><head><title>T</title></head><body>") + for _ in range(2000): # 125 MiB inflated + counter["inflated"] += _CHUNK + # Flushed per chunk, so what is counted is what went on the wire. + yield comp.compress(b"x" * _CHUNK) + comp.flush(zlib.Z_SYNC_FLUSH) + yield comp.flush() + + def handler(request): + return httpx.Response(200, headers={"content-type": "text/html", + "content-encoding": "gzip"}, + content=body()) + async with _client(handler) as c: + meta = await linkpreview.fetch_preview("https://example.com/", client=c) + assert meta is not None + assert counter["inflated"] < 50 * _CHUNK # stopped early, not at 125 MiB + + +async def test_a_declared_oversized_image_is_not_read(resolves_public): + counter = {"sent": 0} + + def handler(request): + return httpx.Response( + 200, headers={"content-type": "image/png", + "content-length": str(linkpreview._MAX_IMAGE_BYTES + 1)}, + content=_endless(counter)) + async with _client(handler) as c: + assert await linkpreview.fetch_image("https://example.com/x.png", client=c) is None + assert counter["sent"] <= _CHUNK diff --git a/packages/meshbay-node/tests/test_roster_pairing.py b/packages/meshbay-node/tests/test_roster_pairing.py index 5c94855..585a909 100644 --- a/packages/meshbay-node/tests/test_roster_pairing.py +++ b/packages/meshbay-node/tests/test_roster_pairing.py @@ -339,6 +339,80 @@ async def test_failures_are_counted_across_connections(tmp_path, roster): "reconnecting must not reset the pairing budget") +async def _grind_wrong_codes(tmp_path, roster, ctx, accounts, per_account): + """Wrong codes from several accounts, each on connections of its own.""" + for n in range(accounts): + uid = f"guesser{n}" + for _ in range(per_account): + session = _session(tmp_path, roster, user_id=uid) + session._ctx = ctx + sk_ed, pk_ed_b64, pk_x_b64 = _keypair() + await session._do_join_request( + _join_msg(session, sk_ed, pk_ed_b64, pk_x_b64, + code="AAAA-AAAA", user_id=uid)) + + +async def test_a_full_code_lock_never_refuses_a_known_member(tmp_path, roster): + """ + Every member reconnecting gets the group key through join_request, so a + lock on wrong codes that also refused *recognised* devices let one member + take the node away from everybody. + """ + GROUP = "b" * 32 + gek = generate_gek() + sk_ed, pk_ed_b64, pk_x_b64, sk_x = _keypair_full() + await roster.pin_identity("grenet", "grenet", pk_ed_b64, pk_x_b64, "code") + await roster.set_member(GROUP, "grenet", ROLE_MEMBER, "active", "cbesson") + + member = _session(tmp_path, roster, user_id="grenet", group_id=GROUP, + gek=gek, join_policy="invite") + await _grind_wrong_codes(tmp_path, roster, member._ctx, accounts=6, + per_account=5) + assert len(member._ctx["join_failures"]) >= 20 + + await member._do_join_request( + _join_msg(member, sk_ed, pk_ed_b64, pk_x_b64, user_id="grenet")) + reply = _last(member) + assert reply.get("ok") is True and reply.get("gek") is True, reply + assert unwrap_gek_aes(reply, *_x_raw(sk_x, pk_x_b64)) == gek + + +async def test_ordinary_refusals_do_not_feed_the_code_lock(tmp_path, roster): + """`code_required` is what every first contact without a code hears.""" + ctx = None + for n in range(30): + session = _session(tmp_path, roster, user_id=f"newcomer{n}") + if ctx is None: + ctx = session._ctx + session._ctx = ctx + sk_ed, pk_ed_b64, pk_x_b64 = _keypair() + await session._do_join_request( + _join_msg(session, sk_ed, pk_ed_b64, pk_x_b64, user_id=f"newcomer{n}")) + assert _last(session).get("reason") == "code_required" + assert not ctx.get("join_failures") + + +async def test_one_account_guessing_locks_itself_not_the_others(tmp_path, roster): + session = _session(tmp_path, roster) + ctx = session._ctx + await _grind_wrong_codes(tmp_path, roster, ctx, accounts=1, per_account=6) + guesser = _session(tmp_path, roster, user_id="guesser0") + guesser._ctx = ctx + sk_ed, pk_ed_b64, pk_x_b64 = _keypair() + await guesser._do_join_request( + _join_msg(guesser, sk_ed, pk_ed_b64, pk_x_b64, code="AAAA-AAAA", + user_id="guesser0")) + assert _last(guesser).get("detail") == "Pairing temporarily locked" + + code = await roster.create_invite("", "grenet", ROLE_OPERATOR, "local-cli") + honest = _session(tmp_path, roster, user_id="grenet") + honest._ctx = ctx + sk_ed, pk_ed_b64, pk_x_b64 = _keypair() + await honest._do_join_request( + _join_msg(honest, sk_ed, pk_ed_b64, pk_x_b64, code=code)) + assert await roster.find_device("grenet", pk_ed_b64) is not None + + async def test_group_id_cannot_name_another_group(tmp_path, roster): session = _session(tmp_path, roster) session._group_id = "a" * 32 |