570 lines
21 KiB
Python
570 lines
21 KiB
Python
#!/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="<f4")
|
|
expected = len(self.works) * self.dimensions
|
|
if self.embeddings.size != expected:
|
|
raise SystemExit(
|
|
f"embedding size mismatch: expected {expected}, got {self.embeddings.size}"
|
|
)
|
|
self.embeddings = self.embeddings.reshape((len(self.works), self.dimensions))
|
|
score_config = self.manifest.get("score") if isinstance(self.manifest, dict) else None
|
|
self.vector_weight = float((score_config or {}).get("vectorWeight", 0.8))
|
|
self.tag_weight = float((score_config or {}).get("tagWeight", 0.2))
|
|
self.query_encoder = self.create_query_encoder(args)
|
|
|
|
@staticmethod
|
|
def _read_json(path: Path) -> 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="<f4")
|
|
|
|
def search(self, payload: dict[str, Any]) -> 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())
|