init
This commit is contained in:
@@ -0,0 +1,569 @@
|
||||
#!/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())
|
||||
Reference in New Issue
Block a user