Spaces:
Running
Running
davanstrien HF Staff
Accept/reject feedback: per-tag votes stored in a private bucket with a pseudonymous id
5605c42 verified Download app.py from davanstrien/hub-task-tagger: direct link, hf CLI and curl.
- Browser
- Download file 23.6 kB
-
https://huggingface.co/spaces/davanstrien/hub-task-tagger/resolve/main/app.py
- Command line
-
hf download hf://spaces/davanstrien/hub-task-tagger/app.py
-
curl -L -o app.py https://huggingface.co/spaces/davanstrien/hub-task-tagger/resolve/main/app.py
23.6 kB
| """Hub task tagger demo: FastAPI backend for the single-page frontend in static/index.html. | |
| The model input is rebuilt exactly as in training (see tagger_core.py). The scoring path, temperature and | |
| threshold are copied from the Gradio test Space (demo-space/app.py), with the base model only. | |
| """ | |
| import asyncio | |
| import hashlib | |
| import hmac | |
| import json | |
| import math | |
| import os | |
| import random | |
| import re | |
| import time | |
| import uuid | |
| from collections import OrderedDict, deque | |
| from contextlib import asynccontextmanager | |
| from datetime import datetime, timezone | |
| from pathlib import Path | |
| import requests | |
| import torch | |
| from fastapi import FastAPI, Request | |
| from fastapi.concurrency import run_in_threadpool | |
| from fastapi.responses import FileResponse, JSONResponse | |
| from gliner2.classification import ( | |
| ClassificationConfig, | |
| ClassificationSchema, | |
| Classifier, | |
| ) | |
| from huggingface_hub import HfApi, hf_hub_download | |
| from huggingface_hub.utils import ( | |
| GatedRepoError, | |
| HfHubHTTPError, | |
| RepositoryNotFoundError, | |
| ) | |
| from pydantic import BaseModel | |
| from tokenizers import Tokenizer | |
| from tagger_core import RETIRED, StateBuilder, fetch_first_rows, render_first_rows | |
| # cpu-basic has 2 vCPUs; os.cpu_count() reports the host, so set the thread count explicitly. | |
| torch.set_num_threads(int(os.environ.get("TORCH_THREADS", "2"))) | |
| HERE = Path(__file__).parent | |
| LABELS = json.loads((HERE / "labels.json").read_text()) | |
| TASK = "labels" # task name in the trained classification schema | |
| MAX_TEXT_CHARS = 2000 # train-gliner2.py --max-text-chars | |
| MAX_STATE_TOKENS = 370 # prepare_hub_tasks.py --max-state-tokens | |
| BASE_TOKENIZER = "Qwen/Qwen3.5-4B-Base" | |
| MODEL_REPO = "davanstrien/hub-task-tagger-gliner2.5-base" | |
| # Temperature fitted on calibration (top-1 log-loss); tau chosen on calibration micro-F1 (threshold rule, top-1 kept). | |
| T, TAU = 1.300, 0.225 | |
| TOP_K = 5 | |
| # Baked into the image at build time (see Dockerfile); falls back to the Hub when run elsewhere. | |
| MODEL_DIR = Path(os.environ.get("MODEL_DIR", "/home/user/models")) | |
| TAGS_URL = "https://huggingface.co/api/datasets-tags-by-type?type=task_categories" | |
| CACHE_SIZE = 2000 | |
| WARM_IDS = ["hubxrt/LPMusicCapsMTT_a2t", "dougdotcon/douvras-failure-atlas", | |
| "mim-chess-vlas/train_800_dense__mask__overlay_a75__sim__agentview_camera__static"] # the page examples | |
| # ---- startup --------------------------------------------------------------------------------------------------- | |
| STARTUP = {} | |
| t0 = time.perf_counter() | |
| if (MODEL_DIR / "tokenizer.json").exists(): | |
| tokenizer_path = MODEL_DIR / "tokenizer.json" | |
| model_path = MODEL_DIR / "base" | |
| schema_path = model_path / "classification_schema.json" | |
| STARTUP["source"] = "image" | |
| else: | |
| tokenizer_path = hf_hub_download(BASE_TOKENIZER, "tokenizer.json") | |
| model_path = MODEL_REPO | |
| schema_path = hf_hub_download(MODEL_REPO, "classification_schema.json") | |
| STARTUP["source"] = "hub" | |
| states = StateBuilder(Tokenizer.from_file(str(tokenizer_path)), MAX_STATE_TOKENS) | |
| STARTUP["tokenizer_load_s"] = round(time.perf_counter() - t0, 2) | |
| MODEL_REVISION = (MODEL_DIR / "revision.txt").read_text().strip() if (MODEL_DIR / "revision.txt").exists() else None | |
| t0 = time.perf_counter() | |
| SCHEMA = ClassificationSchema().multi(TASK, LABELS) | |
| CONFIG = ClassificationConfig(batch_size=1) | |
| CLASSIFIER = Classifier.from_pretrained(str(model_path)).to(device="cpu").eval() | |
| if json.loads(Path(schema_path).read_text())["tasks"][0]["labels"] != LABELS: | |
| raise SystemExit(f"{MODEL_REPO}: trained labels differ from labels.json") | |
| STARTUP["model_load_s"] = round(time.perf_counter() - t0, 2) | |
| STARTUP["torch_threads"] = torch.get_num_threads() | |
| LIVE_TAGS = None | |
| print("startup", STARTUP, flush=True) | |
| # One request at a time on the model: 2 vCPUs do not serve parallel forward passes well. | |
| MODEL_LOCK = asyncio.Lock() | |
| def live_tags(): | |
| global LIVE_TAGS | |
| if LIVE_TAGS is None: | |
| LIVE_TAGS = {x["label"] for x in requests.get(TAGS_URL, timeout=30).json()["task_categories"]} | |
| return LIVE_TAGS | |
| def declared_tags(dataset_id): | |
| """(raw card task_categories, mapped to current tags, gated, sha, error). Anonymous on purpose: public data only.""" | |
| try: | |
| info = HfApi(token=False).dataset_info(dataset_id, expand=["cardData", "gated", "sha"]) | |
| except Exception as error: # noqa: BLE001 -- classified by the caller | |
| return [], [], False, None, error | |
| card = info.card_data.to_dict() if info.card_data else {} | |
| raw = card.get("task_categories") or [] | |
| raw = [raw] if isinstance(raw, str) else [str(t) for t in raw] | |
| live = live_tags() | |
| mapped = sorted({t if t in live else RETIRED[t] for t in raw if t in live or RETIRED.get(t) in live}) | |
| return raw, mapped, bool(info.gated), info.sha, None | |
| def softmax(logits, t): | |
| scaled = [z / t for z in logits] | |
| m = max(scaled) | |
| e = [math.exp(z - m) for z in scaled] | |
| total = sum(e) | |
| return [x / total for x in e] | |
| def score(state): | |
| text = state[:MAX_TEXT_CHARS] | |
| with torch.inference_mode(): | |
| scores = CLASSIFIER.batch_score([text], SCHEMA, config=CONFIG) | |
| logits = scores[0].tasks[TASK] | |
| probs = softmax([logits[label] for label in LABELS], T) | |
| ranked = sorted(zip(LABELS, probs), key=lambda x: -x[1]) | |
| chosen = [label for label, p in ranked if p >= TAU] or [ranked[0][0]] | |
| if ranked[0][0] not in chosen: | |
| chosen.insert(0, ranked[0][0]) | |
| return ranked, chosen | |
| # ---- cache --------------------------------------------------------------------------------------------------- | |
| class PredictionCache: | |
| """LRU of successful responses (without timing), keyed on (dataset_id, revision sha): a new commit misses.""" | |
| def __init__(self, size): | |
| self.size, self.items = size, OrderedDict() | |
| def get(self, dataset_id, sha): | |
| key = (dataset_id, sha) | |
| if sha is None or key not in self.items: | |
| return None | |
| self.items.move_to_end(key) | |
| return self.items[key] | |
| def put(self, dataset_id, sha, body): | |
| if sha is None: | |
| return | |
| self.items[(dataset_id, sha)] = body | |
| self.items.move_to_end((dataset_id, sha)) | |
| while len(self.items) > self.size: | |
| self.items.popitem(last=False) | |
| CACHE = PredictionCache(CACHE_SIZE) | |
| READY = {"warm": False, "warm_s": None, "warmed": []} | |
| class Failure(Exception): | |
| def __init__(self, status, kind, message, dataset_id=None, timing=None): | |
| self.status, self.kind, self.message, self.dataset_id, self.timing = status, kind, message, dataset_id, timing | |
| def clean_id(value): | |
| """Accept `owner/name` or a full dataset URL; return None when it is not a dataset id.""" | |
| value = (value or "").strip().removeprefix("https://huggingface.co/datasets/").split("?")[0].strip("/") | |
| parts = value.split("/") | |
| if len(parts) != 2 or not all(parts) or " " in value: | |
| return None | |
| return value | |
| async def predict_one(raw_id): | |
| """The response body for one dataset id; raises Failure for anything the page should show as an error.""" | |
| dataset_id = clean_id(raw_id) | |
| if dataset_id is None: | |
| raise Failure(400, "bad_id", "That is not a dataset id. Use the form owner/name, for example " | |
| "ar0s/kuka_heat_reaction_expert, or paste the dataset's URL.") | |
| timing = {} | |
| t0 = time.perf_counter() | |
| raw, declared, gated, sha, error = await run_in_threadpool(declared_tags, dataset_id) | |
| timing["card_ms"] = round(1000 * (time.perf_counter() - t0)) | |
| if isinstance(error, GatedRepoError): | |
| gated, error = True, None | |
| if isinstance(error, RepositoryNotFoundError): | |
| raise Failure(404, "not_found", f"No public dataset called {dataset_id}. Check the spelling. " | |
| "Private datasets cannot be read here.", dataset_id, timing) | |
| if error is not None: | |
| status = error.response.status_code if isinstance(error, HfHubHTTPError) and error.response is not None else None | |
| raise Failure(502, "hub_error", f"The Hub did not answer for {dataset_id} (status {status or 'unknown'}). " | |
| "Try again in a minute.", dataset_id, timing) | |
| if gated: | |
| raise Failure(403, "gated", f"{dataset_id} is gated. This page reads only public datasets, " | |
| "so it cannot see the rows. Try a dataset without an access form.", dataset_id, timing) | |
| cached = CACHE.get(dataset_id, sha) | |
| if cached is not None: | |
| return {**cached, "timing": {"cached": True, "card_ms": timing["card_ms"]}} | |
| t0 = time.perf_counter() | |
| fr, cfg, split, fetch_error = await run_in_threadpool(fetch_first_rows, dataset_id) | |
| timing["fetch_ms"] = round(1000 * (time.perf_counter() - t0)) | |
| if fr is None: | |
| print("no preview", dataset_id, fetch_error, flush=True) | |
| raise Failure(422, "no_viewer", f"{dataset_id} has no dataset viewer preview, so there are no rows " | |
| "for the model to read. The viewer is off for some datasets and still processing for new " | |
| "ones. Try again later, or try another dataset.", dataset_id, timing) | |
| t0 = time.perf_counter() | |
| state, truncated = states.build(render_first_rows(fr)) | |
| timing["render_ms"] = round(1000 * (time.perf_counter() - t0), 1) | |
| async with MODEL_LOCK: | |
| t0 = time.perf_counter() | |
| ranked, chosen = await run_in_threadpool(score, state) | |
| timing["model_ms"] = round(1000 * (time.perf_counter() - t0)) | |
| body = { | |
| "dataset_id": dataset_id, | |
| "sha": sha, | |
| "declared": declared, | |
| "declared_raw": raw, | |
| "top": [{"tag": label, "p": round(p, 4)} for label, p in ranked[:TOP_K]], | |
| "set": chosen, | |
| "state": state, | |
| "truncated": truncated, | |
| "preview": f"{cfg}/{split}", | |
| } | |
| CACHE.put(dataset_id, sha, body) | |
| return {**body, "timing": timing} | |
| async def warm_cache(): | |
| """Score the page examples once so the first visitor's clicks are cache hits. Failures are logged, not fatal.""" | |
| t0 = time.perf_counter() | |
| for dataset_id in WARM_IDS: | |
| try: | |
| await predict_one(dataset_id) | |
| READY["warmed"].append(dataset_id) | |
| except Exception as error: # noqa: BLE001 -- warm-up must never stop the app | |
| print("warm-up failed", dataset_id, getattr(error, "kind", repr(error)), flush=True) | |
| READY["warm"], READY["warm_s"] = True, round(time.perf_counter() - t0, 2) | |
| print("warm", READY, flush=True) | |
| # ---- random dataset pool -------------------------------------------------------------------------------------- | |
| POOL_REFRESH_S = 6 * 3600 # rebuild the pool every 6 hours | |
| POOL_RETRY_S = 10 * 60 # retry sooner when a build fails and the pool is empty | |
| POOL_EXPAND = ["gated", "private", "tags", "cardData", "downloads", "likes"] | |
| RANDOM_TRIES = 5 | |
| IS_VALID_URL = "https://datasets-server.huggingface.co/is-valid" | |
| # Any of these words in the id, the tags or the card's tags/pretty_name keeps a dataset out of the pool. | |
| ADULT_WORDS = {"nsfw", "adult", "adults", "porn", "porno", "pornography", "pornographic", "hentai", "xxx", | |
| "sex", "sexual", "sexy", "nude", "nudes", "nudity", "naked", "erotic", "erotica", "explicit", | |
| "r18", "18plus", "lewd", "fetish", "onlyfans", "rule34", "not-for-all-audiences", | |
| # Safety-research sets (jailbreaks, hate speech) are fine to tag by hand, but a random click should | |
| # not put their first row in front of a visitor. | |
| "toxic", "toxicity", "hate", "hateful", "jailbreak", "jailbreaks", "harmful", "offensive", "unsafe", | |
| "redteam", "red-team", "teaming", "abuse", "abusive"} | |
| POOL = {"ids": [], "task_tagged": {}, "size": 0, "built_at": None, "build_s": None, "error": None, | |
| "is_valid_checks": 0, "is_valid_hits": 0} | |
| def words_in(text): | |
| return set(re.split(r"[^a-z0-9+-]+", str(text).lower())) | set(re.split(r"[^a-z0-9]+", str(text).lower())) | |
| def looks_adult(info): | |
| card = info.card_data.to_dict() if info.card_data else {} | |
| card_tags = card.get("tags") or [] | |
| card_tags = [card_tags] if isinstance(card_tags, str) else card_tags | |
| texts = [info.id, card.get("pretty_name") or "", *(info.tags or []), *card_tags] | |
| return any(words_in(text) & ADULT_WORDS for text in texts) | |
| def pool_candidates(api): | |
| """Popular + trending datasets, then recently modified ones with some traction (downloads >= 50 or likes >= 2).""" | |
| yield from api.list_datasets(sort="trendingScore", limit=500, expand=POOL_EXPAND, gated=False) | |
| yield from api.list_datasets(sort="likes", limit=700, expand=POOL_EXPAND, gated=False) | |
| kept = 0 | |
| for info in api.list_datasets(sort="last_modified", limit=20000, expand=POOL_EXPAND, gated=False): | |
| if (info.downloads or 0) >= 50 or (info.likes or 0) >= 2: | |
| kept += 1 | |
| yield info | |
| if kept >= 1000: | |
| break | |
| def build_pool(): | |
| """(ids, {id: has task_categories tag}) of public, ungated, not-adult datasets. Blocking; run in a thread.""" | |
| ids, task_tagged, skipped = [], {}, 0 | |
| for info in pool_candidates(HfApi(token=False)): | |
| if info.id in task_tagged: | |
| continue | |
| if info.private or info.gated or looks_adult(info): | |
| skipped += 1 | |
| continue | |
| ids.append(info.id) | |
| task_tagged[info.id] = any(tag.startswith("task_categories:") for tag in info.tags or []) | |
| return ids, task_tagged, skipped | |
| async def refresh_pool_forever(): | |
| """Build the pool now and every 6 hours, in a thread. Never blocks requests or readiness.""" | |
| while True: | |
| t0 = time.perf_counter() | |
| try: | |
| ids, task_tagged, skipped = await asyncio.to_thread(build_pool) | |
| POOL.update(ids=ids, task_tagged=task_tagged, size=len(ids), built_at=time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), | |
| build_s=round(time.perf_counter() - t0, 1), error=None) | |
| print(f"pool built: {len(ids)} datasets ({skipped} skipped) in {POOL['build_s']} s", flush=True) | |
| except Exception as error: # noqa: BLE001 -- keep the old pool, try again later | |
| POOL["error"] = repr(error) | |
| print("pool build failed", repr(error), flush=True) | |
| await asyncio.sleep(POOL_REFRESH_S if POOL["ids"] else POOL_RETRY_S) | |
| def has_preview(dataset_id): | |
| """True when the dataset viewer has a first-rows preview for this dataset.""" | |
| try: | |
| r = requests.get(IS_VALID_URL, params={"dataset": dataset_id}, timeout=10) | |
| return r.ok and r.json().get("preview") is True | |
| except (requests.RequestException, ValueError): | |
| return False | |
| async def lifespan(app): | |
| # Warm in the background: the port opens at once, and /healthz reports 503 until warm-up ends. | |
| # The random-dataset pool builds in the background too and does not affect /healthz. | |
| tasks = [asyncio.create_task(warm_cache()), asyncio.create_task(refresh_pool_forever())] | |
| yield | |
| for task in tasks: | |
| task.cancel() | |
| # ---- feedback ------------------------------------------------------------------------------------------------ | |
| # A Hub bucket mounted read-write on the Space. One JSON file per event; the app never reads the folder back. | |
| FEEDBACK_DIR = Path(os.environ.get("FEEDBACK_DIR", "/data/feedback")) | |
| FEEDBACK_PER_HOUR = 30 | |
| MODEL_INFO = {"repo": MODEL_REPO, "revision": MODEL_REVISION, "T": T, "tau": TAU} | |
| TAG_RE = re.compile(r"^[a-z0-9][a-z0-9-]{0,63}$") | |
| SHA_RE = re.compile(r"^[0-9a-f]{40}$") | |
| # Rate limiting keys on a salted hash of the client IP. The salt lives only in memory, and neither the IP nor | |
| # the hash is written anywhere. | |
| RATE_SALT = os.urandom(16) | |
| RECENT = {} # hashed ip -> deque of submit times (monotonic seconds) | |
| # Pseudonymous source id: HMAC of month + client IP with a secret salt from the Space secrets. The month in the | |
| # message rotates the ids, so events cannot be linked across months. The IP itself is never stored. | |
| FEEDBACK_SALT = os.environ.get("FEEDBACK_SALT", "").encode() | |
| if not FEEDBACK_SALT: | |
| print("WARNING: FEEDBACK_SALT is not set; feedback events get source_id=null", flush=True) | |
| def check_feedback_dir(): | |
| """True when FEEDBACK_DIR exists and a file can be written and removed there.""" | |
| probe = FEEDBACK_DIR / f".probe-{uuid.uuid4().hex}" | |
| try: | |
| probe.write_text("ok") | |
| probe.unlink() | |
| return True | |
| except OSError as error: | |
| print(f"WARNING: feedback disabled, {FEEDBACK_DIR} is not writable: {error!r}", flush=True) | |
| return False | |
| FEEDBACK = {"enabled": check_feedback_dir(), "dir": str(FEEDBACK_DIR), "written": 0, "write_errors": 0, | |
| "source_ids": bool(FEEDBACK_SALT)} | |
| def source_id(ip, period): | |
| """First 16 hex characters of HMAC-SHA256(FEEDBACK_SALT, "YYYY-MM|ip"), or None without a salt.""" | |
| if not FEEDBACK_SALT: | |
| return None | |
| return hmac.new(FEEDBACK_SALT, f"{period}|{ip}".encode(), hashlib.sha256).hexdigest()[:16] | |
| print("feedback", FEEDBACK, flush=True) | |
| def rate_limited(ip): | |
| key = hashlib.sha256(RATE_SALT + ip.encode()).hexdigest() | |
| now = time.monotonic() | |
| times = RECENT.setdefault(key, deque()) | |
| while times and now - times[0] > 3600: | |
| times.popleft() | |
| if len(times) >= FEEDBACK_PER_HOUR: | |
| return True | |
| times.append(now) | |
| return False | |
| def client_ip(request): | |
| """The address the Space proxy saw. It appends that address last to X-Forwarded-For; earlier entries come from | |
| the client and can be forged, so they are ignored.""" | |
| forwarded = [x.strip() for x in request.headers.get("x-forwarded-for", "").split(",") if x.strip()] | |
| return forwarded[-1] if forwarded else (request.client.host if request.client else "unknown") | |
| class Suggestion(BaseModel): | |
| tag: str | |
| p: float | |
| class FeedbackRequest(BaseModel): | |
| dataset_id: str | |
| sha: str | None = None | |
| model_repo: str | None = None | |
| model_revision: str | None = None | |
| T: float | None = None | |
| tau: float | None = None | |
| suggested: list[Suggestion] | |
| set: list[str] | |
| accepted: list[str] = [] | |
| rejected: list[str] = [] | |
| owner_tags: list[str] = [] | |
| ts: str | None = None | |
| def feedback_problem(fb): | |
| """A short reason when the event is not valid, else None.""" | |
| if clean_id(fb.dataset_id) != fb.dataset_id: | |
| return "dataset_id is not an owner/name id" | |
| if fb.sha is not None and not SHA_RE.match(fb.sha): | |
| return "sha is not a commit sha" | |
| suggested = [s.tag for s in fb.suggested] | |
| if not 1 <= len(suggested) <= TOP_K or len(set(suggested)) != len(suggested): | |
| return f"suggested must list 1 to {TOP_K} different tags" | |
| if not all(tag in LABELS for tag in suggested + fb.set + fb.accepted + fb.rejected): | |
| return "a tag is not one of the model's labels" | |
| if not all(0.0 <= s.p <= 1.0 for s in fb.suggested): | |
| return "a probability is outside 0-1" | |
| if not fb.accepted and not fb.rejected: | |
| return "nothing accepted or rejected" | |
| if set(fb.accepted) & set(fb.rejected): | |
| return "a tag is both accepted and rejected" | |
| if not set(fb.accepted + fb.rejected + fb.set) <= set(suggested): | |
| return "accepted, rejected and set tags must come from the suggested tags" | |
| if len(fb.owner_tags) > 30 or not all(TAG_RE.match(tag) for tag in fb.owner_tags): | |
| return "owner_tags are not valid tags" | |
| if fb.ts is not None and len(fb.ts) > 40: | |
| return "ts is too long" | |
| return None | |
| def write_event(record): | |
| """Write one event as its own file: a temp name in the same folder, then os.replace.""" | |
| day = FEEDBACK_DIR / record["received_at"][:10] | |
| day.mkdir(parents=True, exist_ok=True) | |
| final = day / f"{record['id']}.json" | |
| tmp = day / f".{record['id']}.json.tmp" | |
| tmp.write_text(json.dumps(record, ensure_ascii=False, indent=1)) | |
| os.replace(tmp, final) | |
| # ---- API ------------------------------------------------------------------------------------------------------- | |
| app = FastAPI(title="Hub task tagger demo", docs_url=None, redoc_url=None, lifespan=lifespan) | |
| class PredictRequest(BaseModel): | |
| dataset_id: str | |
| def healthz(): | |
| pool = {k: v for k, v in POOL.items() if k not in ("ids", "task_tagged")} | |
| body = {"ok": READY["warm"], "startup": STARTUP, "warm": READY, "cache_entries": len(CACHE.items), "pool": pool, | |
| "feedback": {k: v for k, v in FEEDBACK.items() if k != "dir"}} | |
| return JSONResponse(status_code=200 if READY["warm"] else 503, content=body) | |
| async def predict(req: PredictRequest): | |
| try: | |
| body = await predict_one(req.dataset_id) | |
| return {**body, "model": MODEL_INFO, "feedback": FEEDBACK["enabled"]} | |
| except Failure as f: | |
| return JSONResponse( | |
| status_code=f.status, | |
| content={"error": f.kind, "message": f.message, "dataset_id": f.dataset_id, "timing": f.timing or {}}, | |
| ) | |
| async def feedback(fb: FeedbackRequest, request: Request): | |
| if not FEEDBACK["enabled"]: | |
| return JSONResponse(status_code=503, content={"error": "feedback_off", "message": "Feedback is switched off."}) | |
| problem = feedback_problem(fb) | |
| if problem: | |
| return JSONResponse(status_code=400, content={"error": "invalid", "message": f"Feedback not saved: {problem}."}) | |
| ip = client_ip(request) | |
| if rate_limited(ip): | |
| return JSONResponse(status_code=429, content={"error": "rate_limited", "message": "Too much feedback from " | |
| "this address in the last hour. Thank you, and please try again later."}) | |
| now = datetime.now(timezone.utc) | |
| period = now.strftime("%Y-%m") | |
| record = { | |
| "id": uuid.uuid4().hex, | |
| "received_at": now.isoformat(timespec="seconds"), | |
| "source_id": source_id(ip, period), | |
| "source_period": period, | |
| **fb.model_dump(), | |
| # The server's own values win over what the page sent. | |
| "model_repo": MODEL_REPO, "model_revision": MODEL_REVISION, "T": T, "tau": TAU, | |
| } | |
| try: | |
| await run_in_threadpool(write_event, record) | |
| except OSError as error: | |
| FEEDBACK["write_errors"] += 1 | |
| print("feedback write failed", repr(error), flush=True) | |
| return JSONResponse(status_code=503, content={"error": "write_failed", "message": "Could not save the " | |
| "feedback just now. Please try again later."}) | |
| FEEDBACK["written"] += 1 | |
| return {"ok": True, "id": record["id"]} | |
| async def random_dataset(): | |
| """A random public dataset from the pool that has a viewer preview. The page then calls /api/predict with it.""" | |
| if not POOL["ids"]: | |
| return JSONResponse(status_code=503, content={"error": "no_pool", "message": "The list of random datasets " | |
| "is not ready yet. Try one of the examples, or try again in a few minutes."}) | |
| for _ in range(RANDOM_TRIES): | |
| dataset_id = random.choice(POOL["ids"]) | |
| ok = await run_in_threadpool(has_preview, dataset_id) | |
| POOL["is_valid_checks"] += 1 | |
| POOL["is_valid_hits"] += ok | |
| if ok: | |
| return {"dataset_id": dataset_id, "has_task_tags": POOL["task_tagged"][dataset_id]} | |
| return JSONResponse(status_code=503, content={"error": "no_preview", "message": "Could not find a random " | |
| "dataset with a viewer preview just now. Try again, or try one of the examples."}) | |
| def index(): | |
| return FileResponse(HERE / "static" / "index.html") | |