import base64 import hashlib import json import os import secrets import sqlite3 import threading import time import zlib from http.cookies import SimpleCookie from pathlib import Path from urllib.parse import urlencode from urllib.request import Request, urlopen def _json_blob(value): if value is None: return None raw = json.dumps(value, ensure_ascii=False, separators=(",", ":")).encode("utf-8") return sqlite3.Binary(zlib.compress(raw, level=6)) def _blob_json(value): if value is None: return None return json.loads(zlib.decompress(value).decode("utf-8")) def _payload_summary(payload): payload = payload or {} workflow = payload.get("workflow", "analysis") mode = payload.get("mode", "tube") model = payload.get("model") or {} strands = payload.get("strands") or [] if workflow == "design": complexes = payload.get("design_complexes") or payload.get("design_targets") or [] tube_sizes = [int(row.get("max_size", 0) or 0) for row in (payload.get("design_tubes") or [])] max_size = max(tube_sizes, default=int((payload.get("design") or {}).get("off_target_max_size", 0) or 0)) else: complexes = [line for line in str(payload.get("complexes_text") or "").splitlines() if line.strip()] max_size = int((payload.get("tube") or {}).get("max_size", 0) or 0) return { "workflow": workflow, "mode": mode, "material": model.get("material", "rna"), "celsius": float(model.get("celsius", 37) or 37), "sodium": float(model.get("sodium", 0) or 0), "magnesium": float(model.get("magnesium", 0) or 0), "max_size": max_size, "strand_count": len(strands), "complex_count": len(complexes), "compute": list(payload.get("compute") or (["design"] if workflow == "design" else [])), "trials": int((payload.get("design") or {}).get("trials", 0) or 0), "stop_condition": float((payload.get("design") or {}).get("stop_condition", 0) or 0), "max_time_seconds": int((payload.get("design") or {}).get("max_time_seconds", 0) or 0), } class AccountStore: def __init__(self, path): self.path = Path(path) self._init_lock = threading.Lock() self._initialized = False def _connect(self): self.path.parent.mkdir(parents=True, exist_ok=True) connection = sqlite3.connect(self.path, timeout=30) connection.row_factory = sqlite3.Row connection.execute("PRAGMA busy_timeout = 30000") connection.execute("PRAGMA foreign_keys = ON") return connection def initialize(self): with self._init_lock: if self._initialized: return with self._connect() as connection: connection.execute("PRAGMA journal_mode = WAL") connection.execute("PRAGMA synchronous = NORMAL") connection.executescript( """ CREATE TABLE IF NOT EXISTS users ( user_id TEXT PRIMARY KEY, username TEXT NOT NULL, email TEXT, display_name TEXT, groups_json TEXT NOT NULL DEFAULT '[]', created_at REAL NOT NULL, last_seen_at REAL NOT NULL ); CREATE TABLE IF NOT EXISTS jobs ( job_id TEXT PRIMARY KEY, user_id TEXT NOT NULL REFERENCES users(user_id) ON DELETE CASCADE, status TEXT NOT NULL, source TEXT NOT NULL DEFAULT 'job', created_at REAL NOT NULL, updated_at REAL NOT NULL, elapsed_seconds REAL, workflow TEXT NOT NULL, mode TEXT NOT NULL, material TEXT NOT NULL, celsius REAL NOT NULL, sodium REAL NOT NULL, magnesium REAL NOT NULL, max_size INTEGER NOT NULL DEFAULT 0, strand_count INTEGER NOT NULL DEFAULT 0, complex_count INTEGER NOT NULL DEFAULT 0, compute_json TEXT NOT NULL DEFAULT '[]', trials INTEGER NOT NULL DEFAULT 0, stop_condition REAL NOT NULL DEFAULT 0, max_time_seconds INTEGER NOT NULL DEFAULT 0, payload_blob BLOB NOT NULL, result_blob BLOB, error_blob BLOB, stored_bytes INTEGER NOT NULL DEFAULT 0 ); CREATE INDEX IF NOT EXISTS jobs_user_created_idx ON jobs(user_id, created_at DESC); CREATE INDEX IF NOT EXISTS jobs_user_status_idx ON jobs(user_id, status, created_at DESC); CREATE TABLE IF NOT EXISTS shares ( share_id TEXT PRIMARY KEY, job_id TEXT NOT NULL REFERENCES jobs(job_id) ON DELETE CASCADE, user_id TEXT NOT NULL REFERENCES users(user_id) ON DELETE CASCADE, active INTEGER NOT NULL DEFAULT 1, created_at REAL NOT NULL, expires_at REAL, last_access_at REAL, access_count INTEGER NOT NULL DEFAULT 0 ); CREATE INDEX IF NOT EXISTS shares_user_created_idx ON shares(user_id, created_at DESC); """ ) self._initialized = True def upsert_user(self, user): self.initialize() now = time.time() with self._connect() as connection: connection.execute( """ INSERT INTO users(user_id, username, email, display_name, groups_json, created_at, last_seen_at) VALUES(?, ?, ?, ?, ?, ?, ?) ON CONFLICT(user_id) DO UPDATE SET username=excluded.username, email=excluded.email, display_name=excluded.display_name, groups_json=excluded.groups_json, last_seen_at=excluded.last_seen_at """, ( user["user_id"], user.get("username") or user["user_id"], user.get("email"), user.get("display_name"), json.dumps(user.get("groups") or [], ensure_ascii=False), now, now, ), ) def create_job(self, job_id, user, payload, status="queued", source="job", created_at=None, result=None, error=None): self.upsert_user(user) created_at = float(created_at or time.time()) summary = _payload_summary(payload) payload_blob = _json_blob(payload) result_blob = _json_blob(result) error_blob = _json_blob(error) stored_bytes = sum(len(item) for item in (payload_blob, result_blob, error_blob) if item is not None) with self._connect() as connection: connection.execute( """ INSERT OR IGNORE INTO jobs( job_id, user_id, status, source, created_at, updated_at, workflow, mode, material, celsius, sodium, magnesium, max_size, strand_count, complex_count, compute_json, trials, stop_condition, max_time_seconds, payload_blob, result_blob, error_blob, stored_bytes ) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) """, ( job_id, user["user_id"], status, source, created_at, created_at, summary["workflow"], summary["mode"], summary["material"], summary["celsius"], summary["sodium"], summary["magnesium"], summary["max_size"], summary["strand_count"], summary["complex_count"], json.dumps(summary["compute"]), summary["trials"], summary["stop_condition"], summary["max_time_seconds"], payload_blob, result_blob, error_blob, stored_bytes, ), ) def update_job(self, job_id, status, *, result=None, error=None, elapsed_seconds=None): self.initialize() fields = ["status = ?", "updated_at = ?"] values = [status, time.time()] if result is not None: fields.append("result_blob = ?") values.append(_json_blob(result)) if error is not None: fields.append("error_blob = ?") values.append(_json_blob(error)) if elapsed_seconds is not None: fields.append("elapsed_seconds = ?") values.append(float(elapsed_seconds)) values.append(job_id) with self._connect() as connection: connection.execute(f"UPDATE jobs SET {', '.join(fields)} WHERE job_id = ?", values) connection.execute( "UPDATE jobs SET stored_bytes = length(payload_blob) + coalesce(length(result_blob), 0) + coalesce(length(error_blob), 0) WHERE job_id = ?", (job_id,), ) def owner_id(self, job_id): self.initialize() with self._connect() as connection: row = connection.execute("SELECT user_id FROM jobs WHERE job_id = ?", (job_id,)).fetchone() return row["user_id"] if row else None def get_job(self, user_id, job_id, include_content=True): self.initialize() columns = "*" if include_content else "job_id,user_id,status,source,created_at,updated_at,elapsed_seconds,workflow,mode,material,celsius,sodium,magnesium,max_size,strand_count,complex_count,compute_json,trials,stop_condition,max_time_seconds,stored_bytes" with self._connect() as connection: row = connection.execute(f"SELECT {columns} FROM jobs WHERE job_id = ? AND user_id = ?", (job_id, user_id)).fetchone() return self._job_row(row, include_content=include_content) if row else None def list_jobs(self, user_id, filters=None): self.initialize() filters = filters or {} limit = min(100, max(1, int(filters.get("limit", 30)))) offset = max(0, int(filters.get("offset", 0))) clauses = ["user_id = ?"] values = [user_id] for field in ("status", "workflow", "mode", "material"): value = str(filters.get(field) or "").strip() if value: clauses.append(f"{field} = ?") values.append(value) search = str(filters.get("q") or "").strip() if search: clauses.append("(job_id LIKE ? OR workflow LIKE ? OR mode LIKE ? OR material LIKE ?)") values.extend([f"%{search}%"] * 4) where = " AND ".join(clauses) columns = "job_id,user_id,status,source,created_at,updated_at,elapsed_seconds,workflow,mode,material,celsius,sodium,magnesium,max_size,strand_count,complex_count,compute_json,trials,stop_condition,max_time_seconds,stored_bytes" with self._connect() as connection: total = connection.execute(f"SELECT count(*) AS count FROM jobs WHERE {where}", values).fetchone()["count"] rows = connection.execute( f"SELECT {columns} FROM jobs WHERE {where} ORDER BY created_at DESC LIMIT ? OFFSET ?", [*values, limit, offset], ).fetchall() return {"items": [self._job_row(row, include_content=False) for row in rows], "total": total, "limit": limit, "offset": offset} def usage(self, user_id): self.initialize() with self._connect() as connection: totals = connection.execute( """SELECT count(*) AS total_jobs, coalesce(sum(CASE WHEN status='success' THEN 1 ELSE 0 END),0) AS success_jobs, coalesce(sum(CASE WHEN status='error' THEN 1 ELSE 0 END),0) AS error_jobs, coalesce(sum(CASE WHEN status IN ('queued','running','cancel_requested') THEN 1 ELSE 0 END),0) AS active_jobs, coalesce(sum(elapsed_seconds),0) AS compute_seconds, coalesce(sum(stored_bytes),0) AS stored_bytes FROM jobs WHERE user_id=?""", (user_id,), ).fetchone() shares = connection.execute("SELECT count(*) AS count FROM shares WHERE user_id=? AND active=1", (user_id,)).fetchone()["count"] result = dict(totals) result["active_shares"] = shares return result def import_history(self, user, entries): imported = 0 for entry in entries: payload = entry.get("payload") if isinstance(entry, dict) else None if not isinstance(payload, dict): continue result = entry.get("result") if isinstance(entry.get("result"), dict) else None raw_id = str(entry.get("id") or entry.get("job_id") or uuid_token()) job_id = f"import-{hashlib.sha256((user['user_id'] + ':' + raw_id).encode()).hexdigest()[:24]}" created = entry.get("created_at") or entry.get("created_at_iso") try: created_at = float(created) except (TypeError, ValueError): try: created_at = time.mktime(time.strptime(str(created).split(".")[0].replace("Z", ""), "%Y-%m-%dT%H:%M:%S")) except (TypeError, ValueError): created_at = time.time() before = self.get_job(user["user_id"], job_id, include_content=False) self.create_job(job_id, user, payload, status="success" if result else "input_only", source="import", created_at=created_at, result=result) if before is None: imported += 1 return imported def delete_job(self, user_id, job_id): self.initialize() with self._connect() as connection: row = connection.execute("SELECT status FROM jobs WHERE job_id=? AND user_id=?", (job_id, user_id)).fetchone() if row is None: return False if row["status"] in {"queued", "running", "cancel_requested"}: raise ValueError("Active jobs cannot be deleted.") connection.execute("DELETE FROM jobs WHERE job_id=? AND user_id=?", (job_id, user_id)) return True def create_share(self, user_id, job_id, expires_in=None): job = self.get_job(user_id, job_id, include_content=False) if job is None: raise ValueError("Job not found.") now = time.time() expires_at = now + int(expires_in) if expires_in else None share_id = uuid_token(18) with self._connect() as connection: connection.execute( "INSERT INTO shares(share_id,job_id,user_id,active,created_at,expires_at) VALUES(?,?,?,?,?,?)", (share_id, job_id, user_id, 1, now, expires_at), ) return self.get_share_for_owner(user_id, share_id) def share_metadata(self, share_id): self.initialize() now = time.time() with self._connect() as connection: row = connection.execute( """SELECT s.share_id, s.job_id, s.user_id, s.created_at, s.expires_at, s.active FROM shares s WHERE s.share_id=? AND s.active=1 AND (s.expires_at IS NULL OR s.expires_at>?)""", (share_id, now), ).fetchone() return dict(row) if row else None def record_share_access(self, share_id): self.initialize() now = time.time() with self._connect() as connection: connection.execute( "UPDATE shares SET access_count=access_count+1,last_access_at=? WHERE share_id=?", (now, share_id), ) def list_shares(self, user_id): self.initialize() with self._connect() as connection: rows = connection.execute( """SELECT s.*,j.status,j.workflow,j.mode,j.material,j.created_at AS job_created_at FROM shares s JOIN jobs j ON j.job_id=s.job_id WHERE s.user_id=? ORDER BY s.created_at DESC""", (user_id,), ).fetchall() return [dict(row) for row in rows] def get_share_for_owner(self, user_id, share_id): self.initialize() with self._connect() as connection: row = connection.execute("SELECT * FROM shares WHERE share_id=? AND user_id=?", (share_id, user_id)).fetchone() return dict(row) if row else None def update_share(self, user_id, share_id, *, active=None, expires_in="unchanged"): share = self.get_share_for_owner(user_id, share_id) if share is None: return None fields = [] values = [] if active is not None: fields.append("active=?") values.append(1 if active else 0) if expires_in != "unchanged": fields.append("expires_at=?") values.append(time.time() + int(expires_in) if expires_in else None) if fields: values.extend([share_id, user_id]) with self._connect() as connection: connection.execute(f"UPDATE shares SET {', '.join(fields)} WHERE share_id=? AND user_id=?", values) return self.get_share_for_owner(user_id, share_id) def resolve_share(self, share_id, *, record_access=True): self.initialize() now = time.time() with self._connect() as connection: row = connection.execute( """SELECT s.share_id, s.created_at AS share_created_at, j.job_id, j.user_id, j.status, j.source, j.created_at, j.updated_at, j.elapsed_seconds, j.workflow, j.mode, j.material, j.celsius, j.sodium, j.magnesium, j.max_size, j.strand_count, j.complex_count, j.compute_json, j.trials, j.stop_condition, j.max_time_seconds, j.payload_blob, j.result_blob, j.error_blob, j.stored_bytes FROM shares s JOIN jobs j ON j.job_id=s.job_id WHERE s.share_id=? AND s.active=1 AND (s.expires_at IS NULL OR s.expires_at>?)""", (share_id, now), ).fetchone() if row is None: return None if record_access: connection.execute("UPDATE shares SET access_count=access_count+1,last_access_at=? WHERE share_id=?", (now, share_id)) job = self._job_row(row, include_content=True) return { "id": share_id, "share_id": share_id, "job_id": job["job_id"], "created_at": row["share_created_at"], "status": job["status"], "updated_at": job["updated_at"], "elapsed_seconds": job["elapsed_seconds"], "payload": job["payload"], "result": job["result"], "error": job["error"], "result_summary": job["result_summary"], } @staticmethod def _job_row(row, include_content): item = dict(row) item["compute"] = json.loads(item.pop("compute_json") or "[]") item["result_summary"] = { key: item[key] for key in ( "workflow", "mode", "material", "celsius", "sodium", "magnesium", "max_size", "strand_count", "complex_count", "compute", "trials", "stop_condition", "max_time_seconds", ) } if include_content: item["payload"] = _blob_json(item.pop("payload_blob")) item["result"] = _blob_json(item.pop("result_blob")) item["error"] = _blob_json(item.pop("error_blob")) return item def uuid_token(size=24): return secrets.token_urlsafe(size) class OIDCAuth: def __init__(self, redis_factory): self.redis_factory = redis_factory self.issuer = os.environ.get("NP_OIDC_ISSUER", "https://auth.lihato.icu/application/o/nupack-account/").rstrip("/") + "/" self.client_id = os.environ.get("NP_OIDC_CLIENT_ID", "np-replica-web") self.redirect_uri = os.environ.get("NP_OIDC_REDIRECT_URI", "https://np.lihato.icu/auth/callback") self.post_logout_uri = os.environ.get("NP_OIDC_POST_LOGOUT_URI", "https://np.lihato.icu/") self.scopes = os.environ.get("NP_OIDC_SCOPES", "openid profile email").strip() self.cookie_name = os.environ.get("NP_AUTH_COOKIE_NAME", "np_session") self.session_ttl = int(os.environ.get("NP_AUTH_SESSION_TTL_SECONDS", str(7 * 86400))) self.required = os.environ.get("NP_AUTH_REQUIRED", "1") != "0" self._configuration = None self._config_lock = threading.Lock() def configuration(self): with self._config_lock: if self._configuration is None: self._configuration = self._request_json(self.issuer + ".well-known/openid-configuration") return self._configuration @staticmethod def _request_json(url, *, data=None, headers=None): body = urlencode(data).encode("utf-8") if data is not None else None request = Request(url, data=body, headers=headers or {}) with urlopen(request, timeout=15) as response: return json.loads(response.read().decode("utf-8")) def begin_login(self, next_path="/"): if not next_path.startswith("/") or next_path.startswith("//"): next_path = "/" state = uuid_token(24) verifier = uuid_token(48) challenge = base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode() self.redis_factory().setex( f"np_replica:oidc_state:{state}", 600, json.dumps({"verifier": verifier, "next": next_path}), ) config = self.configuration() query = urlencode({ "client_id": self.client_id, "response_type": "code", "redirect_uri": self.redirect_uri, "scope": self.scopes, "state": state, "code_challenge": challenge, "code_challenge_method": "S256", }) return f"{config['authorization_endpoint']}?{query}" def complete_login(self, params): state = (params.get("state") or [""])[0] code = (params.get("code") or [""])[0] if not state or not code: raise ValueError("OIDC callback is missing code or state.") key = f"np_replica:oidc_state:{state}" client = self.redis_factory() raw = client.get(key) client.delete(key) if raw is None: raise ValueError("OIDC login state has expired.") pending = json.loads(raw) config = self.configuration() tokens = self._request_json(config["token_endpoint"], data={ "grant_type": "authorization_code", "client_id": self.client_id, "code": code, "redirect_uri": self.redirect_uri, "code_verifier": pending["verifier"], }, headers={"Content-Type": "application/x-www-form-urlencoded"}) claims = self._request_json(config["userinfo_endpoint"], headers={"Authorization": f"Bearer {tokens['access_token']}"}) subject = str(claims.get("sub") or "").strip() if not subject: raise ValueError("OIDC userinfo did not return a subject.") user = { "user_id": subject, "username": claims.get("preferred_username") or claims.get("nickname") or claims.get("email") or subject, "email": claims.get("email"), "display_name": claims.get("name") or claims.get("preferred_username") or subject, "groups": claims.get("groups") or [], } session_id = uuid_token(32) client.setex( f"np_replica:session:{session_id}", self.session_ttl, json.dumps({"user": user, "id_token": tokens.get("id_token")}, ensure_ascii=False), ) return session_id, user, pending.get("next") or "/" def current_session(self, headers): if not self.required: user = {"user_id": "development", "username": "development", "email": None, "display_name": "Development", "groups": []} return {"user": user, "session_id": None} cookie = SimpleCookie() cookie.load(headers.get("Cookie", "")) morsel = cookie.get(self.cookie_name) if morsel is None: return None session_id = morsel.value raw = self.redis_factory().get(f"np_replica:session:{session_id}") if raw is None: return None session = json.loads(raw) session["session_id"] = session_id return session def logout_url(self, session): config = self.configuration() endpoint = config.get("end_session_endpoint") if not endpoint: return "/" query = {"post_logout_redirect_uri": self.post_logout_uri} if session and session.get("id_token"): query["id_token_hint"] = session["id_token"] return f"{endpoint}?{urlencode(query)}" def delete_session(self, session): if session and session.get("session_id"): self.redis_factory().delete(f"np_replica:session:{session['session_id']}") def cookie_header(self, session_id): return f"{self.cookie_name}={session_id}; Path=/; Max-Age={self.session_ttl}; HttpOnly; Secure; SameSite=Lax" def clear_cookie_header(self): return f"{self.cookie_name}=; Path=/; Max-Age=0; HttpOnly; Secure; SameSite=Lax"