#!/usr/bin/env python3 """Serve the WebUI and provide local natural-language vector search.""" from __future__ import annotations import argparse import http.client import json import math import os from functools import partial from http import HTTPStatus from http.server import SimpleHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from typing import Any from urllib.error import HTTPError, URLError from urllib.parse import parse_qs, urlparse from urllib.request import Request, urlopen DEFAULT_HOST = "127.0.0.1" DEFAULT_PORT = 8000 def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Serve built Astro WebUI and expose /api/search for natural-language ASMR search." ) parser.add_argument("--host", default=DEFAULT_HOST, help="Bind host.") parser.add_argument("--port", type=int, default=DEFAULT_PORT, help="Bind port.") parser.add_argument("--web-dir", default="web/dist", help="Directory to serve as static WebUI.") parser.add_argument("--data-dir", default="web/data", help="Directory containing manifest/works/embeddings.") parser.add_argument("--batch-size", type=int, default=1, help="SentenceTransformer encode batch size.") parser.add_argument( "--api-key-env", default="EMBEDDING_API_KEY", help="Environment variable that contains the remote embeddings API key.", ) parser.add_argument( "--remote-base-url", default=None, help="Override manifest baseUrl for remote-openai-compatible search.", ) parser.add_argument( "--remote-query-prefix", default="", help="Optional prefix prepended to remote search queries, for example 'query: '.", ) parser.add_argument("--timeout", type=float, default=120.0, help="Remote request timeout seconds.") return parser.parse_args() def normalize(value: Any) -> str: return str(value or "").lower().strip() def parse_limit(value: Any, default: int = 48, maximum: int = 200) -> int: return max(1, min(maximum, int(value or default))) def searchable_text(work: dict[str, Any]) -> str: values = [ work.get("title"), work.get("sourceId"), work.get("circle"), *(work.get("vas") or []), *(work.get("tags") or []), ] return " ".join(normalize(value) for value in values) def recent_sort_value(work: dict[str, Any]) -> str: return str(work.get("createDate") or work.get("release") or work.get("sourceId") or "") def embeddings_url(base_url: str) -> str: return f"{base_url.rstrip('/')}/v1/embeddings" def normalize_vector(vector: list[float]) -> list[float]: norm = math.sqrt(sum(value * value for value in vector)) if norm == 0: return vector return [value / norm for value in vector] def request_remote_embedding( *, url: str, model: str, query: str, api_key: str, timeout: float, ) -> list[float]: body = json.dumps({"model": model, "input": [query]}).encode("utf-8") headers = {"Content-Type": "application/json"} if api_key: headers["Authorization"] = f"Bearer {api_key}" request = Request(url, data=body, headers=headers, method="POST") try: with urlopen(request, timeout=timeout) as response: payload = json.load(response) except HTTPError as exc: detail = exc.read(1000).decode("utf-8", errors="replace") raise RuntimeError(f"remote embeddings API returned HTTP {exc.code}: {detail}") from exc except ( URLError, TimeoutError, json.JSONDecodeError, http.client.IncompleteRead, http.client.RemoteDisconnected, ConnectionResetError, OSError, ) as exc: raise RuntimeError(f"remote embeddings request failed: {exc}") from exc data = payload.get("data") if isinstance(payload, dict) else None if not isinstance(data, list) or not data: raise RuntimeError("remote embeddings response does not contain data[]") first = data[0] if not isinstance(first, dict): raise RuntimeError("remote embeddings response data[0] is not an object") embedding = first.get("embedding") if not isinstance(embedding, list) or not embedding: raise RuntimeError("remote embeddings response does not contain embedding") try: return [float(value) for value in embedding] except (TypeError, ValueError) as exc: raise RuntimeError("remote embedding contains non-numeric values") from exc class LocalSentenceTransformerQueryEncoder: def __init__(self, model_name: str, batch_size: int) -> None: try: from sentence_transformers import SentenceTransformer except ImportError as exc: raise SystemExit( "Local natural-language search requires semantic dependencies. " "Run `uv run --extra semantic python scripts/search_server.py`." ) from exc self.model_name = model_name self.batch_size = batch_size self.model = SentenceTransformer(model_name) def encode(self, query: str) -> list[float]: text = f"query: {query}" if "e5" in self.model_name.lower() else query embedding = self.model.encode( [text], batch_size=self.batch_size, normalize_embeddings=True, show_progress_bar=False, )[0] return [float(value) for value in embedding] class RemoteOpenAICompatibleQueryEncoder: def __init__( self, *, base_url: str, model: str, api_key: str, timeout: float, query_prefix: str, should_normalize: bool, ) -> None: self.url = embeddings_url(base_url) self.model = model self.api_key = api_key self.timeout = timeout self.query_prefix = query_prefix self.should_normalize = should_normalize def encode(self, query: str) -> list[float]: vector = request_remote_embedding( url=self.url, model=self.model, query=f"{self.query_prefix}{query}", api_key=self.api_key, timeout=self.timeout, ) return normalize_vector(vector) if self.should_normalize else vector class SearchIndex: def __init__(self, args: argparse.Namespace, data_dir: Path) -> None: try: import numpy as np except ImportError as exc: raise SystemExit( "Natural-language search requires numpy. " "Run `uv run --extra remote python scripts/search_server.py` or " "`uv run --extra semantic python scripts/search_server.py`." ) from exc self.np = np self.manifest = self._read_json(data_dir / "manifest.json") self.works = self._read_json(data_dir / "works.json") self.method = str(self.manifest.get("method") or "") self.dimensions = int(self.manifest["dimensions"]) self.model_name = str(self.manifest["model"]) embeddings_path = data_dir / str(self.manifest["embeddingFile"]) self.embeddings = np.fromfile(embeddings_path, dtype=" Any: with path.open(encoding="utf-8") as file: return json.load(file) def create_query_encoder(self, args: argparse.Namespace): if self.method == "sentence-transformers": return LocalSentenceTransformerQueryEncoder(self.model_name, args.batch_size) if self.method == "remote-openai-compatible": base_url = args.remote_base_url or self.manifest.get("baseUrl") if not base_url: raise SystemExit( "remote-openai-compatible manifest does not contain baseUrl. " "Pass --remote-base-url." ) return RemoteOpenAICompatibleQueryEncoder( base_url=str(base_url), model=self.model_name, api_key=os.environ.get(args.api_key_env, ""), timeout=args.timeout, query_prefix=args.remote_query_prefix, should_normalize=bool(self.manifest.get("normalized", True)), ) raise SystemExit( f"unsupported embedding method for natural-language search: {self.method}. " "Use sentence-transformers or remote-openai-compatible data." ) def query_embedding(self, query: str): embedding = self.query_encoder.encode(query) if len(embedding) != self.dimensions: raise RuntimeError( f"query embedding dimensions mismatch: expected {self.dimensions}, got {len(embedding)}" ) return self.np.asarray(embedding, dtype=" dict[str, Any]: query = str(payload.get("query") or "").strip() if not query: raise ValueError("query is required") limit = parse_limit(payload.get("limit")) ages = {normalize(value) for value in payload.get("ages") or [] if normalize(value)} required_tags = [normalize(value) for value in payload.get("requiredTags") or [] if normalize(value)] try: query_vector = self.query_embedding(query) except RuntimeError as exc: raise ValueError(str(exc)) from exc scores = self.embeddings @ query_vector candidates = self.filtered_indices(ages, required_tags) if candidates.size == 0: return {"results": []} candidate_scores = scores[candidates] take = min(limit, candidate_scores.size) top_positions = self.np.argpartition(-candidate_scores, take - 1)[:take] top_positions = top_positions[self.np.argsort(-candidate_scores[top_positions])] results = [ { "index": int(candidates[position]), "score": float(candidate_scores[position]), "work": self.works[int(candidates[position])], } for position in top_positions ] return {"results": results} def list_works(self, payload: dict[str, Any]) -> dict[str, Any]: limit = parse_limit(payload.get("limit"), default=100) query = str(payload.get("query") or payload.get("q") or "").strip() ages = {normalize(value) for value in payload.get("ages") or [] if normalize(value)} required_tags = [ normalize(value) for value in (payload.get("requiredTags") or payload.get("tags") or []) if normalize(value) ] entries = [ work for work in self.works if self.matches_work(work, ages=ages, required_tags=required_tags, query=query) ] entries.sort(key=recent_sort_value, reverse=True) return {"results": [{"work": work} for work in entries[:limit]]} def get_work(self, identifier: Any) -> dict[str, Any]: normalized_id = normalize(identifier) if not normalized_id: raise ValueError("id is required") for work in self.works: if ( normalize(work.get("sourceId")) == normalized_id or normalize(work.get("id")) == normalized_id or normalize(work.get("embeddingIndex")) == normalized_id ): return {"work": work} raise ValueError("work not found") def similar(self, payload: dict[str, Any]) -> dict[str, Any]: try: source_index = int(payload.get("index")) except (TypeError, ValueError) as exc: raise ValueError("index is required") from exc if not 0 <= source_index < len(self.works): raise ValueError("index is out of range") limit = parse_limit(payload.get("limit")) filters = payload.get("filters") if isinstance(payload.get("filters"), dict) else {} ages = {normalize(value) for value in filters.get("ages") or [] if normalize(value)} exclude_same_group = bool(filters.get("excludeSameGroup")) source_work = self.works[source_index] source_group = source_work.get("groupId") source_tags = set(source_work.get("tags") or []) scores = self.embeddings @ self.embeddings[source_index] candidates: list[int] = [] combined_scores: list[float] = [] vector_scores: list[float] = [] tag_scores: list[float] = [] for index, work in enumerate(self.works): if index == source_index: continue if exclude_same_group and source_group and source_group == work.get("groupId"): continue if ages and normalize(work.get("ageCategory") or "unknown") not in ages: continue vector_score = float(scores[index]) tag_score = tag_jaccard(source_tags, work.get("tags") or []) score = self.vector_weight * vector_score + self.tag_weight * tag_score candidates.append(index) combined_scores.append(score) vector_scores.append(vector_score) tag_scores.append(tag_score) if not candidates: return {"results": []} order = sorted(range(len(candidates)), key=lambda position: combined_scores[position], reverse=True) results = [ { "index": int(candidates[position]), "score": float(combined_scores[position]), "vector": float(vector_scores[position]), "tag": float(tag_scores[position]), "work": self.works[int(candidates[position])], } for position in order[:limit] ] return {"results": results} def filtered_indices(self, ages: set[str], required_tags: list[str]): indices: list[int] = [] for index, work in enumerate(self.works): if ages and normalize(work.get("ageCategory")) not in ages: continue if required_tags: tags = " ".join(normalize(tag) for tag in work.get("tags") or []) if not all(tag in tags for tag in required_tags): continue indices.append(index) return self.np.asarray(indices, dtype="int64") @staticmethod def matches_work( work: dict[str, Any], *, ages: set[str], required_tags: list[str], query: str, ) -> bool: if ages and normalize(work.get("ageCategory") or "unknown") not in ages: return False if required_tags: tags = " ".join(normalize(tag) for tag in work.get("tags") or []) if not all(tag in tags for tag in required_tags): return False if query and normalize(query) not in searchable_text(work): return False return True def tag_jaccard(source_tags: set[str], candidate_tags: list[Any]) -> float: if not source_tags and not candidate_tags: return 0.0 intersection = 0 seen: set[Any] = set() for tag in candidate_tags: if tag in seen: continue seen.add(tag) if tag in source_tags: intersection += 1 union = len(source_tags) + len(seen) - intersection return intersection / union if union else 0.0 class SearchRequestHandler(SimpleHTTPRequestHandler): search_index: SearchIndex def end_headers(self) -> None: self.send_header("Access-Control-Allow-Origin", "*") self.send_header("Access-Control-Allow-Headers", "content-type") self.send_header("Access-Control-Allow-Methods", "GET, POST, OPTIONS") super().end_headers() def do_OPTIONS(self) -> None: self.send_response(HTTPStatus.NO_CONTENT) self.end_headers() def do_GET(self) -> None: parsed = urlparse(self.path) if parsed.path == "/api/status": self.write_json( { "ok": True, "count": len(self.search_index.works), "model": self.search_index.model_name, } ) return if parsed.path == "/api/search": params = parse_qs(parsed.query) payload = { "query": (params.get("q") or [""])[0], "limit": (params.get("limit") or [48])[0], } self.handle_search(payload) return if parsed.path == "/api/works": params = parse_qs(parsed.query) payload = { "query": (params.get("q") or [""])[0], "limit": (params.get("limit") or [100])[0], "ages": params.get("age") or [], "requiredTags": params.get("tag") or [], } self.handle_works(payload) return if parsed.path == "/api/work": params = parse_qs(parsed.query) self.handle_work((params.get("id") or [""])[0]) return if parsed.path == "/api/similar": params = parse_qs(parsed.query) payload = { "index": (params.get("index") or [None])[0], "limit": (params.get("limit") or [48])[0], "filters": { "ages": params.get("age") or [], "excludeSameGroup": (params.get("excludeSameGroup") or [""])[0] in {"1", "true"}, }, } self.handle_similar(payload) return super().do_GET() def do_POST(self) -> None: parsed = urlparse(self.path) if parsed.path not in {"/api/search", "/api/similar", "/api/works", "/api/work"}: self.send_error(HTTPStatus.NOT_FOUND) return try: length = int(self.headers.get("content-length") or 0) body = self.rfile.read(length).decode("utf-8") if length else "{}" payload = json.loads(body) if not isinstance(payload, dict): raise ValueError("request body must be a JSON object") except (UnicodeDecodeError, json.JSONDecodeError, ValueError) as exc: self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return if parsed.path == "/api/search": self.handle_search(payload) elif parsed.path == "/api/similar": self.handle_similar(payload) elif parsed.path == "/api/works": self.handle_works(payload) else: self.handle_work(payload.get("id")) def handle_search(self, payload: dict[str, Any]) -> None: try: self.write_json(self.search_index.search(payload)) except ValueError as exc: self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def handle_works(self, payload: dict[str, Any]) -> None: try: self.write_json(self.search_index.list_works(payload)) except ValueError as exc: self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def handle_work(self, identifier: Any) -> None: try: self.write_json(self.search_index.get_work(identifier)) except ValueError as exc: status = HTTPStatus.NOT_FOUND if str(exc) == "work not found" else HTTPStatus.BAD_REQUEST self.write_json({"error": str(exc)}, status) def handle_similar(self, payload: dict[str, Any]) -> None: try: self.write_json(self.search_index.similar(payload)) except ValueError as exc: self.write_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def write_json(self, payload: dict[str, Any], status: HTTPStatus = HTTPStatus.OK) -> None: body = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8") self.send_response(status) self.send_header("Content-Type", "application/json; charset=utf-8") self.send_header("Content-Length", str(len(body))) self.end_headers() self.wfile.write(body) def main() -> int: args = parse_args() web_dir = Path(args.web_dir).resolve() data_dir = Path(args.data_dir).resolve() if not web_dir.exists(): raise SystemExit(f"web directory not found: {web_dir}") if not data_dir.exists(): raise SystemExit(f"data directory not found: {data_dir}") print("Loading natural-language search index...") SearchRequestHandler.search_index = SearchIndex(args, data_dir) handler = partial(SearchRequestHandler, directory=str(web_dir)) server = ThreadingHTTPServer((args.host, args.port), handler) print(f"Serving http://{args.host}:{args.port}") print(f"Search model: {SearchRequestHandler.search_index.model_name}") print(f"Search method: {SearchRequestHandler.search_index.method}") try: server.serve_forever() except KeyboardInterrupt: print("\nShutting down...") finally: server.server_close() return 0 if __name__ == "__main__": raise SystemExit(main())