From 30cdd059140ea1ef6ce445ba04445dd22658015a Mon Sep 17 00:00:00 2001 From: Lihatoo <1747565629@gmail.com> Date: Sun, 26 Jul 2026 03:00:54 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E8=AE=BE=E8=AE=A1=E8=AE=A1?= =?UTF-8?q?=E7=AE=97=E9=94=99=E8=AF=AF=E5=B9=B6=E4=BC=98=E5=8C=96=E5=88=86?= =?UTF-8?q?=E4=BA=AB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docker-compose.yml | 4 +- service/account.py | 539 ++++++++++++++++++++++++++++++++++++++++ service/index.html | 65 ++++- service/server.py | 238 ++++++++++++++++-- service/test_account.py | 98 ++++++++ 5 files changed, 916 insertions(+), 28 deletions(-) create mode 100644 service/account.py create mode 100644 service/test_account.py diff --git a/docker-compose.yml b/docker-compose.yml index b1bd892..3232cf3 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -9,7 +9,7 @@ services: - ./runtime/redis:/data np-replica: - image: np-app:v4.7 + image: 100.64.0.2:13000/lihato/np_app:v4.8 build: context: . dockerfile: Dockerfile @@ -45,7 +45,7 @@ services: GOTO_NUM_THREADS: 4 np-worker: - image: np-app:v4.7 + image: 100.64.0.2:13000/lihato/np_app:v4.8 container_name: np-worker restart: always mem_limit: 56g diff --git a/service/account.py b/service/account.py new file mode 100644 index 0000000..29ac634 --- /dev/null +++ b/service/account.py @@ -0,0 +1,539 @@ +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" diff --git a/service/index.html b/service/index.html index da3e825..d6257f4 100644 --- a/service/index.html +++ b/service/index.html @@ -1703,6 +1703,13 @@ A+B +
@@ -2221,6 +2228,9 @@ A+B design_wobble_label: "Wobble mutations", design_wobble_allow: "允许", design_wobble_prohibit: "禁止", + design_fixed_target_policy: "固定目标策略", + design_fixed_target_exclude: "剥离出优化", + design_fixed_target_include: "参与优化", design_tab_targets: "Target Tubes", design_tab_hard: "Hard Constraints", design_tab_soft: "Soft Constraints", @@ -2440,6 +2450,7 @@ A+B design_ensemble_defect: "Ensemble defect", design_trials_used: "试验数", design_seed: "最佳种子", + design_fixed_excluded: "固定目标剥离", design_constraint: "约束", design_definition: "输入定义", design_sequence: "设计序列", @@ -2568,6 +2579,9 @@ A+B design_wobble_label: "Wobble Mutations", design_wobble_allow: "Allow", design_wobble_prohibit: "Prohibit", + design_fixed_target_policy: "Fixed Target Policy", + design_fixed_target_exclude: "Exclude from Optimization", + design_fixed_target_include: "Include in Optimization", design_tab_targets: "Target Tubes", design_tab_hard: "Hard Constraints", design_tab_soft: "Soft Constraints", @@ -2787,6 +2801,7 @@ A+B design_ensemble_defect: "Ensemble Defect", design_trials_used: "Trials", design_seed: "Best Seed", + design_fixed_excluded: "Fixed Targets Excluded", design_constraint: "Constraint", design_definition: "Input Definition", design_sequence: "Designed Sequence", @@ -3803,9 +3818,14 @@ A+B const result = entry?.result && typeof entry.result === "object" ? entry.result : null; return { id: String(entry.id || entry.job_id || `${Date.now()}-${index}-${Math.random().toString(16).slice(2)}`), + job_id: entry.job_id || null, + status: entry.status || (result ? "success" : "input_only"), created_at: entry.created_at_iso || entry.created_at || new Date().toISOString(), + updated_at: entry.updated_at || null, + elapsed_seconds: entry.elapsed_seconds ?? null, payload, result, + error: entry.error || null, result_summary: buildHistorySummary(payload, entry.result_summary || {}), }; } @@ -4050,6 +4070,7 @@ A+B document.getElementById("designSeed").value = payload.design?.seed ?? 0; document.getElementById("designWobble").value = String(payload.design?.wobble_mutations ?? false); document.getElementById("designMaxTimeHours").value = ((payload.design?.max_time_seconds ?? 0) / 3600); + document.getElementById("designFixedTargetPolicy").value = payload.design?.fixed_target_policy ?? "exclude_from_optimization"; const payloadCompute = Array.isArray(payload.compute) ? payload.compute : []; document.querySelectorAll(".check input").forEach((input) => { input.checked = payloadCompute.includes(input.value); @@ -4094,28 +4115,50 @@ A+B } else { lastPayload = payload; lastResult = null; - renderEmptyState(); + if (["queued", "running", "cancel_requested"].includes(item.status)) { + const label = item.status === "queued" + ? t("status_queued") + : (item.status === "cancel_requested" ? t("status_canceling") : t("status_polling")); + results.innerHTML = `

${escapeHtml(label)}

${t("history_job_id")}: ${escapeHtml(item.job_id || item.id || "-")}

`; + setStatus(formatStatus(label)); + } else { + renderEmptyState(); + } } } + function sharedJobIsActive(item) { + return ["queued", "running", "cancel_requested"].includes(item?.status); + } + + async function fetchSharedItem(shareId, { poll = false } = {}) { + const suffix = poll ? "?poll=1" : ""; + const { response, data } = await fetchJsonOrThrow(`/api/shares/${encodeURIComponent(shareId)}${suffix}`); + if (!response.ok || data.status !== "success") { + throw new Error(data?.error || `Share request failed with HTTP ${response.status}`); + } + const item = normalizeHistoryEntry(data.share, 0); + if (!item) throw new Error("Shared record is invalid."); + return item; + } + async function loadSharedHistoryFromUrl() { const pathMatch = window.location.pathname.match(/^\/share\/([A-Za-z0-9_-]{8,128})$/); const shareId = pathMatch?.[1] || new URLSearchParams(window.location.search).get("share"); if (!shareId) return; document.body.classList.add("share-mode"); try { - const { response, data } = await fetchJsonOrThrow(`/api/shares/${encodeURIComponent(shareId)}`); - if (!response.ok || data.status !== "success") { - throw new Error(data?.error || `Share request failed with HTTP ${response.status}`); - } - const item = normalizeHistoryEntry(data.share, 0); - if (!item) throw new Error("Shared record is invalid."); - item.error = data.share.error || null; + let item = await fetchSharedItem(shareId); applyHistoryItem(item); document.querySelectorAll(".control-panel input, .control-panel select, .control-panel textarea, .control-panel button").forEach((node) => { if (!node.closest(".toolbar") && !node.matches("[data-design-tab]")) node.disabled = true; }); setStatus(formatStatus(t("history_shared_loaded"))); + while (sharedJobIsActive(item)) { + await new Promise((resolve) => setTimeout(resolve, 1500)); + item = await fetchSharedItem(shareId, { poll: true }); + applyHistoryItem(item); + } } catch (error) { setStatus(`${t("history_share_failed")}: ${error.message}`, true); } @@ -4133,7 +4176,7 @@ A+B
- ${["success", "error", "input_only", "canceled"].includes(item.status) ? `` : ""} + ${["queued", "running", "cancel_requested"].includes(item.status) ? "" : ``}
@@ -4329,6 +4372,7 @@ A+B seed: Number(document.getElementById("designSeed").value), wobble_mutations: document.getElementById("designWobble").value === "true", max_time_seconds: Math.round(Number(document.getElementById("designMaxTimeHours").value || 0) * 3600), + fixed_target_policy: document.getElementById("designFixedTargetPolicy").value, }, hard_constraints: getHardConstraints(), soft_constraints: getSoftConstraints(), @@ -4519,6 +4563,7 @@ A+B function renderDesignSummaryCard(result) { const stats = result.design?.stats || {}; + const excludedFixedCount = result.design?.optimization?.excluded_fixed_targets?.length || 0; return `

${t("design_summary_card")}

@@ -4526,6 +4571,7 @@ A+B
${t("design_ensemble_defect")}${Number(result.design?.ensemble_defect || 0).toFixed(6)}
${t("design_trials_used")}${result.options?.trials ?? 1}
${t("design_seed")}${stats.seed ?? "-"}
+
${t("design_fixed_excluded")}${excludedFixedCount}
${t("summary_elapsed")}${stats.design_time ? `${Number(stats.design_time).toFixed(3)} s` : "-"}
@@ -6130,6 +6176,7 @@ A+B document.getElementById("designSeed").value = example.design?.seed ?? 0; document.getElementById("designWobble").value = String(example.design?.wobble_mutations ?? false); document.getElementById("designMaxTimeHours").value = ((example.design?.max_time_seconds ?? 0) / 3600); + document.getElementById("designFixedTargetPolicy").value = example.design?.fixed_target_policy ?? "exclude_from_optimization"; document.getElementById("complexesText").value = example.complexes_text; const exampleCompute = Array.isArray(example.compute) ? example.compute : []; document.querySelectorAll(".check input").forEach((input) => { diff --git a/service/server.py b/service/server.py index f939cf5..fefc7bf 100644 --- a/service/server.py +++ b/service/server.py @@ -82,8 +82,11 @@ UNIT_SCALE = { } IUPAC_CODES = "ACGTUWSMKRYBDHVN" IUPAC_CONSTRAINT_TOKEN = re.compile(rf"[{IUPAC_CODES}](?:\d+)?") +FIXED_BASES = set("ACGTU") +COMPLEMENT_BASE = str.maketrans({"A": "T", "C": "G", "G": "C", "T": "A", "U": "A"}) VALID_COMPUTE = {"pfunc", "pairs", "mfe", "sample", "subopt", "ensemble_size"} +FIXED_TARGET_POLICIES = {"exclude_from_optimization", "include"} CANCEL_REQUESTED_STATUS = "cancel_requested" CANCELED_STATUS = "canceled" TERMINAL_JOB_STATUSES = {"success", "error", CANCELED_STATUS} @@ -245,6 +248,54 @@ def is_valid_iupac_constraint(sequence): return True +def expand_iupac_constraint(sequence): + seq = normalize_design_sequence(sequence) + output = [] + index = 0 + while index < len(seq): + match = IUPAC_CONSTRAINT_TOKEN.match(seq, index) + if not match: + raise ValueError(f"Invalid sequence constraint near: {seq[index:]}") + token = match.group(0) + base = token[0] + count = int(token[1:] or "1") + output.append(base * count) + index = match.end() + return "".join(output) + + +def is_mutable_iupac_constraint(sequence): + expanded = expand_iupac_constraint(sequence) + return any(base not in FIXED_BASES for base in expanded) + + +def reverse_complement_fixed(sequence): + return sequence.upper().translate(COMPLEMENT_BASE)[::-1] + + +def parse_domain_tokens(text): + tokens = [token.strip() for token in re.split(r"[\s,]+", str(text or "").strip()) if token.strip()] + output = [] + for token in tokens: + complement = False + domain_name = token + if token.startswith("~"): + complement = True + domain_name = token[1:] + elif token.endswith("*"): + complement = True + domain_name = token[:-1] + output.append((domain_name, complement)) + return output + + +def get_mapping_value(mapping, key, default=None): + try: + return mapping[key] + except Exception: + return default + + def ensure_unit_interval_limits(lower, upper, label): if not (0 <= lower <= upper <= 1): raise ValueError(f"{label} limits must satisfy 0 <= lower <= upper <= 1.") @@ -320,6 +371,7 @@ def parse_design_options(payload): "seed": int(raw.get("seed", 0)), "wobble_mutations": bool(raw.get("wobble_mutations", False)), "max_time_seconds": int(raw.get("max_time_seconds", 0)), + "fixed_target_policy": str(raw.get("fixed_target_policy", "exclude_from_optimization")).strip().lower(), } if options["trials"] < 1 or options["trials"] > 8: raise ValueError("design trials must be between 1 and 8.") @@ -331,6 +383,8 @@ def parse_design_options(payload): raise ValueError("design stop_condition must be between 0 and 1.") if options["max_time_seconds"] < 0: raise ValueError("design max_time_seconds must be non-negative.") + if options["fixed_target_policy"] not in FIXED_TARGET_POLICIES: + raise ValueError("design fixed_target_policy must be exclude_from_optimization or include.") return options @@ -388,11 +442,14 @@ def build_design_domains(domain_payload): domain = Domain(sequence_constraint, name=name) domain_map[name] = domain + mutable = is_mutable_iupac_constraint(sequence_constraint) ordered.append( { "name": name, "constraint": sequence_constraint, "object": domain, + "mutable": mutable, + "fixed_sequence": None if mutable else expand_iupac_constraint(sequence_constraint), } ) return domain_map, ordered @@ -424,11 +481,12 @@ def parse_domain_composition(text, domain_map): return domains -def build_design_strands(strand_payload, domain_map=None): +def build_design_strands(strand_payload, domain_map=None, domain_rows=None): if not strand_payload: raise ValueError("At least one design strand is required.") domain_map = domain_map or {} + domain_info = {item["name"]: item for item in (domain_rows or [])} use_domain_composition = bool(domain_map) target_strand_map = {} ordered = [] @@ -457,11 +515,28 @@ def build_design_strands(strand_payload, domain_map=None): target_strand = TargetStrand(strand_domains, name=name) constraint_kind = "sequence_constraint" constraint_value = sequence_constraint + mutable = is_mutable_iupac_constraint(sequence_constraint) + fixed_sequence = None if mutable else expand_iupac_constraint(sequence_constraint) else: inline_domain = None target_strand = TargetStrand(strand_domains, name=name) constraint_kind = "domain_composition" constraint_value = raw_definition + fixed_parts = [] + mutable = False + for domain_name, complement in parse_domain_tokens(raw_definition): + info = domain_info.get(domain_name) + if info is None: + mutable = True + fixed_parts = [] + break + if info.get("mutable"): + mutable = True + fixed_parts = [] + break + part = info.get("fixed_sequence") or "" + fixed_parts.append(reverse_complement_fixed(part) if complement else part) + fixed_sequence = None if mutable else "".join(fixed_parts) else: if not is_valid_iupac_constraint(sequence_constraint): raise ValueError(f"Design strand {name} contains unsupported constraint characters.") @@ -470,6 +545,8 @@ def build_design_strands(strand_payload, domain_map=None): target_strand = TargetStrand(strand_domains, name=name) constraint_kind = "sequence_constraint" constraint_value = sequence_constraint + mutable = is_mutable_iupac_constraint(sequence_constraint) + fixed_sequence = None if mutable else expand_iupac_constraint(sequence_constraint) target_strand_map[name] = target_strand ordered.append( @@ -481,6 +558,8 @@ def build_design_strands(strand_payload, domain_map=None): "object": target_strand, "domains": strand_domains, "domain": inline_domain, + "mutable": mutable, + "fixed_sequence": fixed_sequence, } ) @@ -728,13 +807,14 @@ def build_soft_constraints(payload, domain_map, strand_map, target_complex_map): return constraints -def parse_design_complexes(payload, target_strand_map): +def parse_design_complexes(payload, target_strand_map, strand_rows=None): target_rows = payload.get("design_complexes") or payload.get("design_targets") or [] if not target_rows: raise ValueError("At least one design target complex is required.") targets = [] target_complex_map = {} + strand_info = {item["name"]: item for item in (strand_rows or [])} for idx, row in enumerate(target_rows, start=1): name = (row.get("name") or f"target_{idx}").strip() or f"target_{idx}" @@ -761,11 +841,14 @@ def parse_design_complexes(payload, target_strand_map): structure, name=name, ) + mutable = any(strand_info.get(token, {}).get("mutable", True) for token in tokens) target_payload = { "name": name, "strands": tokens, "structure": structure, "object": target_complex, + "mutable": mutable, + "optimization_status": "included", } targets.append(target_payload) target_complex_map[name] = target_complex @@ -859,6 +942,21 @@ def parse_design_tubes(payload, target_rows, target_complex_map, default_max_siz return ordered_rows, tubes +def payload_with_allowed_design_targets(payload, allowed_target_names): + allowed = set(allowed_target_names) + filtered_tubes = [] + for row in payload.get("design_tubes") or []: + on_targets = [entry for entry in (row.get("on_targets") or []) if (entry.get("complex") or "").strip() in allowed] + if on_targets: + tube_row = dict(row) + tube_row["on_targets"] = on_targets + filtered_tubes.append(tube_row) + + output = dict(payload) + output["design_tubes"] = filtered_tubes + return output + + def validate_design_object_names(design_domains, design_strands, target_rows, tube_rows): name_map = {} for kind, rows in ( @@ -1206,46 +1304,62 @@ def serialize_design_result( ordered_domains, ordered_strands, ): + analysis_map = getattr(design_result, "to_analysis", {}) or {} designed_domains = [] designed_domain_map = getattr(design_result, "domains", {}) or {} for item in ordered_domains: designed_domain = designed_domain_map.get(item["object"]) + domain_sequence = str(designed_domain) if designed_domain is not None else item.get("fixed_sequence") designed_domains.append( { "name": item["name"], "constraint": item["constraint"], - "sequence": str(designed_domain) if designed_domain is not None else None, - "length": len(str(designed_domain)) if designed_domain is not None else None, + "sequence": domain_sequence, + "length": len(domain_sequence) if domain_sequence is not None else None, + "mutable": bool(item.get("mutable", True)), } ) designed_strands = [] + designed_strand_by_name = {} for item in ordered_strands: target_strand = item["object"] - analysis_strand = design_result.to_analysis[target_strand] + analysis_strand = get_mapping_value(analysis_map, target_strand) + sequence = str(analysis_strand) if analysis_strand is not None else item.get("fixed_sequence") designed_strands.append( { "name": item["name"], "constraint": item["constraint"], "constraint_kind": item.get("constraint_kind", "sequence_constraint"), "definition": item.get("definition", item["constraint"]), - "sequence": str(analysis_strand), - "length": len(str(analysis_strand)), + "sequence": sequence, + "length": len(sequence) if sequence is not None else None, + "mutable": bool(item.get("mutable", True)), } ) + if sequence is not None: + designed_strand_by_name[item["name"]] = sequence target_complexes = [] for target in target_rows: target_complex = target["object"] - analysis_complex = design_result.to_analysis[target_complex] + analysis_complex = get_mapping_value(analysis_map, target_complex) + if analysis_complex is not None: + display = stringify_complex(analysis_complex) + sequence = flatten_sequence(analysis_complex) + else: + display = " + ".join(target["strands"]) + sequence = "".join(designed_strand_by_name.get(name, "") for name in target["strands"]) target_complexes.append( { "name": target["name"], - "display": stringify_complex(analysis_complex), + "display": display, "strand_names": list(target["strands"]), "structure": target["structure"], - "sequence": flatten_sequence(analysis_complex), + "sequence": sequence, "target_concentration_M": target.get("target_concentration_M"), + "optimization_status": target.get("optimization_status", "included"), + "mutable": bool(target.get("mutable", True)), } ) @@ -1291,6 +1405,11 @@ def serialize_design_result( "seed": design_options["seed"], "wobble_mutations": design_options["wobble_mutations"], "max_time_seconds": design_options["max_time_seconds"], + "fixed_target_policy": design_options["fixed_target_policy"], + }, + "optimization": { + "included_targets": [row["name"] for row in target_rows if row.get("optimization_status") == "included"], + "excluded_fixed_targets": [row["name"] for row in target_rows if row.get("optimization_status") == "fixed_excluded"], }, "target_tubes": [ { @@ -1346,22 +1465,42 @@ def run_job_payload(payload): target_strand_map, design_strands = build_design_strands( payload.get("strands") or [], domain_map=design_domain_map, + domain_rows=design_domains, ) - target_rows, target_complex_map = parse_design_complexes(payload, target_strand_map) + target_rows, target_complex_map = parse_design_complexes(payload, target_strand_map, strand_rows=design_strands) + optimization_target_rows = target_rows + optimization_target_complex_map = target_complex_map + if design_options["fixed_target_policy"] == "exclude_from_optimization": + optimization_target_rows = [row for row in target_rows if row.get("mutable", True)] + fixed_names = {row["name"] for row in target_rows if not row.get("mutable", True)} + for row in target_rows: + if row["name"] in fixed_names: + row["optimization_status"] = "fixed_excluded" + optimization_target_complex_map = { + row["name"]: row["object"] for row in optimization_target_rows + } + if not optimization_target_rows: + raise ValueError("Design contains no mutable target complexes after fixed-target filtering.") hard_constraints = build_hard_constraints(payload, design_domain_map, target_strand_map) soft_constraints = build_soft_constraints( payload, design_domain_map, target_strand_map, - target_complex_map, + optimization_target_complex_map, ) tube_rows = [] design_tubes = [] if mode == "tube": + optimization_payload = payload + if design_options["fixed_target_policy"] == "exclude_from_optimization": + optimization_payload = payload_with_allowed_design_targets( + payload, + {row["name"] for row in optimization_target_rows}, + ) tube_rows, design_tubes = parse_design_tubes( - payload, - target_rows, - target_complex_map, + optimization_payload, + optimization_target_rows, + optimization_target_complex_map, design_options["off_target_max_size"], ) validate_design_object_names(design_domains, design_strands, target_rows, tube_rows) @@ -1386,7 +1525,7 @@ def run_job_payload(payload): ) elif mode == "complex": design_job = complex_design( - complexes=[row["object"] for row in target_rows], + complexes=[row["object"] for row in optimization_target_rows], model=model, options=design_job_options, hard_constraints=hard_constraints, @@ -1923,6 +2062,69 @@ def get_share(share_id): return dict(item) if item else None +def public_share_payload(share_id, *, record_access=True): + metadata = ACCOUNT_STORE.share_metadata(share_id) + if metadata is None: + return get_share(share_id) + + live = get_job_data(metadata["job_id"], include_payload=True) + if live is not None: + if record_access: + ACCOUNT_STORE.record_share_access(share_id) + payload = live.get("payload") + if payload is None: + account_job = ACCOUNT_STORE.get_job(metadata["user_id"], metadata["job_id"], include_content=True) + payload = account_job.get("payload") if account_job else None + result = live.get("result") + error = live.get("error") + return { + "id": share_id, + "share_id": share_id, + "job_id": metadata["job_id"], + "created_at": metadata["created_at"], + "status": live.get("status"), + "updated_at": live.get("updated_at"), + "elapsed_seconds": live.get("elapsed_seconds"), + "payload": payload, + "result": result, + "error": error, + "result_summary": build_history_summary(payload, live), + } + + return ACCOUNT_STORE.resolve_share(share_id, record_access=record_access) + + +def build_history_summary(payload, fallback=None): + payload = payload or {} + fallback = fallback or {} + model = payload.get("model") or {} + workflow = payload.get("workflow") or fallback.get("workflow") or "analysis" + mode = payload.get("mode") or fallback.get("mode") or "tube" + 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", fallback.get("material", "rna")), + "celsius": float(model.get("celsius", fallback.get("celsius", 37)) or 37), + "sodium": float(model.get("sodium", fallback.get("sodium", 0)) or 0), + "magnesium": float(model.get("magnesium", fallback.get("magnesium", 0)) or 0), + "max_size": max_size, + "strand_count": len(strands), + "complex_count": len(complexes), + "compute": list(payload.get("compute") or fallback.get("compute") or ([] if workflow != "design" else ["design"])), + "trials": int((payload.get("design") or {}).get("trials", fallback.get("trials", 0)) or 0), + "stop_condition": float((payload.get("design") or {}).get("stop_condition", fallback.get("stop_condition", 0)) or 0), + "max_time_seconds": int((payload.get("design") or {}).get("max_time_seconds", fallback.get("max_time_seconds", 0)) or 0), + } + + def recover_interrupted_jobs(): client = redis_client() recovered = 0 @@ -2043,6 +2245,7 @@ EXAMPLE_PAYLOAD = { "seed": 0, "wobble_mutations": False, "max_time_seconds": 0, + "fixed_target_policy": "exclude_from_optimization", }, "design_domains": [ {"name": "a", "sequence": "N10"}, @@ -2104,6 +2307,7 @@ DESIGN_TUBE_EXAMPLE_PAYLOAD = { "seed": 1, "wobble_mutations": False, "max_time_seconds": 0, + "fixed_target_policy": "exclude_from_optimization", }, "design_domains": [ {"name": "a", "sequence": "N10"}, @@ -2329,7 +2533,7 @@ class AppHandler(BaseHTTPRequestHandler): if parsed.path.startswith("/api/shares/"): share_id = parsed.path.rsplit("/", 1)[-1] - share = ACCOUNT_STORE.resolve_share(share_id) or get_share(share_id) + share = public_share_payload(share_id, record_access=(parse_qs(parsed.query).get("poll") or ["0"])[0] != "1") if share is None: self._respond(*json_bytes({"error": "Share not found"}, status=HTTPStatus.NOT_FOUND)) return diff --git a/service/test_account.py b/service/test_account.py new file mode 100644 index 0000000..e4d0e1b --- /dev/null +++ b/service/test_account.py @@ -0,0 +1,98 @@ +import tempfile +import time +import unittest +from pathlib import Path + +from service.account import AccountStore + + +def user(user_id): + return { + "user_id": user_id, + "username": user_id, + "display_name": user_id.title(), + "email": f"{user_id}@example.test", + "groups": [], + } + + +class AccountStoreTest(unittest.TestCase): + def setUp(self): + self.temporary_directory = tempfile.TemporaryDirectory() + self.store = AccountStore(Path(self.temporary_directory.name) / "account.sqlite3") + self.alice = user("alice") + self.bob = user("bob") + self.payload = { + "workflow": "analysis", + "mode": "tube", + "model": {"material": "rna", "celsius": 37, "sodium": 1, "magnesium": 0}, + "strands": [{"name": "A", "sequence": "ACGU"}], + "complexes_text": "A", + "tube": {"max_size": 1}, + "compute": ["mfe"], + } + + def tearDown(self): + self.temporary_directory.cleanup() + + def test_jobs_are_private_and_content_round_trips(self): + result = {"workflow": "analysis", "complex_results": [{"name": "A"}]} + self.store.create_job("job-1", self.alice, self.payload) + self.store.update_job("job-1", "success", result=result, elapsed_seconds=1.25) + + self.assertIsNone(self.store.get_job(self.bob["user_id"], "job-1")) + item = self.store.get_job(self.alice["user_id"], "job-1") + self.assertEqual(item["payload"], self.payload) + self.assertEqual(item["result"], result) + self.assertEqual(item["elapsed_seconds"], 1.25) + self.assertEqual(self.store.list_jobs("alice")["total"], 1) + self.assertEqual(self.store.list_jobs("bob")["total"], 0) + + def test_share_can_expire_and_be_disabled(self): + self.store.create_job("job-2", self.alice, self.payload, status="success", result={"ok": True}) + share = self.store.create_share("alice", "job-2") + resolved = self.store.resolve_share(share["share_id"]) + self.assertEqual(resolved["job_id"], "job-2") + self.assertEqual(resolved["result"], {"ok": True}) + + self.store.update_share("alice", share["share_id"], active=False) + self.assertIsNone(self.store.resolve_share(share["share_id"])) + self.assertIsNone(self.store.update_share("bob", share["share_id"], active=True)) + + expiring = self.store.create_share("alice", "job-2", expires_in=1) + with self.store._connect() as connection: + connection.execute( + "UPDATE shares SET expires_at=? WHERE share_id=?", + (time.time() - 1, expiring["share_id"]), + ) + self.assertIsNone(self.store.resolve_share(expiring["share_id"])) + + def test_running_job_can_be_shared_and_later_exposes_result(self): + self.store.create_job("job-running", self.alice, self.payload, status="running") + share = self.store.create_share("alice", "job-running") + + resolved = self.store.resolve_share(share["share_id"]) + self.assertEqual(resolved["job_id"], "job-running") + self.assertEqual(resolved["status"], "running") + self.assertIsNone(resolved["result"]) + + self.store.update_job("job-running", "success", result={"ok": True}, elapsed_seconds=2.5) + resolved = self.store.resolve_share(share["share_id"]) + self.assertEqual(resolved["status"], "success") + self.assertEqual(resolved["result"], {"ok": True}) + self.assertEqual(resolved["elapsed_seconds"], 2.5) + + self.assertIsNone(self.store.update_share("bob", share["share_id"], active=False)) + self.store.update_share("alice", share["share_id"], active=False) + self.assertIsNone(self.store.resolve_share(share["share_id"])) + + def test_active_job_cannot_be_deleted(self): + self.store.create_job("job-3", self.alice, self.payload) + with self.assertRaises(ValueError): + self.store.delete_job("alice", "job-3") + self.store.update_job("job-3", "canceled", error={"message": "canceled"}) + self.assertTrue(self.store.delete_job("alice", "job-3")) + + +if __name__ == "__main__": + unittest.main()