Skip to content

Commit d33dbc0

Browse files
committed
UPDATE: Harmonize Dockerfile and FastAPI server with mcp implementation, add TINYSEARCH_VERSION environment variable and normalize research query function
1 parent 646a0e4 commit d33dbc0

5 files changed

Lines changed: 107 additions & 95 deletions

File tree

Dockerfile

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@ LABEL org.opencontainers.image.title="TinySearch" \
1111
ENV PYTHONDONTWRITEBYTECODE=1 \
1212
PYTHONUNBUFFERED=1 \
1313
PIP_NO_CACHE_DIR=1 \
14+
TINYSEARCH_VERSION=${TINYSEARCH_VERSION} \
1415
TINYSEARCH_MODELS_DIR=/data/models \
1516
PLAYWRIGHT_BROWSERS_PATH=/ms-playwright
1617

servers/fastapi_server.py

Lines changed: 20 additions & 84 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from __future__ import annotations
22

33
import asyncio
4+
import os
45
import sys
56
from contextlib import asynccontextmanager
67
from pathlib import Path
@@ -18,6 +19,8 @@
1819
from services.research_config_service import (
1920
config_trace_path,
2021
load_research_config,
22+
normalize_research_query,
23+
research_run_kwargs,
2124
research_tokenizer_name,
2225
)
2326
from services.site_crawl_service import crawl_search
@@ -36,6 +39,10 @@ async def _ensure_local_bundle_for_config(config: dict[str, Any]) -> None:
3639
await asyncio.to_thread(ensure_onnx_bundle_sync, str(config["embedding_model"]))
3740

3841

42+
def _tinysearch_version() -> str:
43+
return os.environ.get("TINYSEARCH_VERSION", "dev").strip() or "dev"
44+
45+
3946
@asynccontextmanager
4047
async def _lifespan(_app: FastAPI):
4148
cfg = load_research_config()
@@ -46,7 +53,7 @@ async def _lifespan(_app: FastAPI):
4653
app = FastAPI(
4754
title="TinySearch API",
4855
description="Web search, site crawl, and hybrid research endpoints.",
49-
version="0.1.4",
56+
version=_tinysearch_version(),
5057
lifespan=_lifespan,
5158
)
5259

@@ -85,6 +92,7 @@ class ResearchRequest(BaseModel):
8592
chunk_max_per_source_url: int | None = Field(default=None, ge=0, le=500)
8693
max_concurrent_crawls: int | None = Field(default=None, ge=1, le=20)
8794
max_concurrent_embedding_calls: int | None = Field(default=None, ge=1, le=20)
95+
pipeline_timeout_seconds: float | None = Field(default=None, gt=0)
8896
embedding_timeout_seconds: float | None = Field(default=None, gt=0)
8997
embedding_timeout_retries: int | None = Field(default=None, ge=0, le=10)
9098
crawl_fit_markdown_mode: str | None = None
@@ -173,94 +181,22 @@ async def site_crawl_get(
173181
@app.post("/research")
174182
async def research_endpoint(request: ResearchRequest) -> dict[str, Any]:
175183
config = load_research_config()
176-
embedding_model = request.embedding_model or str(config["embedding_model"])
184+
query = normalize_research_query(request.query)
185+
overrides = request.model_dump(exclude_none=True)
186+
overrides.pop("query")
187+
trace_path = overrides.pop("trace_path", None)
188+
189+
run_kwargs = research_run_kwargs(config)
190+
run_kwargs.update(overrides)
191+
embedding_model = str(run_kwargs["embedding_model"])
177192
if normalize_embedding_backend(str(config["embedding_backend"])) == "onnx":
178193
from services.onnx_bundle_service import ensure_onnx_bundle_sync
179194

180195
await asyncio.to_thread(ensure_onnx_bundle_sync, embedding_model)
181196
result = await agentic_run(
182-
request.query,
183-
search_top_k=request.search_top_k or int(config["search_top_k"]),
184-
search_rrf_cutoff=request.search_rrf_cutoff
185-
if request.search_rrf_cutoff is not None
186-
else float(config["search_rrf_cutoff"]),
187-
search_dense_weight=request.search_dense_weight
188-
if request.search_dense_weight is not None
189-
else float(config["search_dense_weight"]),
190-
search_max_results_to_keep=request.search_max_results_to_keep
191-
or int(config["search_max_results_to_keep"]),
192-
chunk_rrf_cutoff=request.chunk_rrf_cutoff
193-
if request.chunk_rrf_cutoff is not None
194-
else float(config["chunk_rrf_cutoff"]),
195-
chunk_dense_weight=request.chunk_dense_weight
196-
if request.chunk_dense_weight is not None
197-
else float(config["chunk_dense_weight"]),
198-
chunk_max_results_to_keep=request.chunk_max_results_to_keep
199-
or int(config["chunk_max_results_to_keep"]),
200-
chunk_rank_oversample=request.chunk_rank_oversample
201-
or int(config["chunk_rank_oversample"]),
202-
chunk_dedupe_jaccard_threshold=request.chunk_dedupe_jaccard_threshold
203-
if request.chunk_dedupe_jaccard_threshold is not None
204-
else float(config["chunk_dedupe_jaccard_threshold"]),
205-
chunk_max_per_source_url=request.chunk_max_per_source_url
206-
if request.chunk_max_per_source_url is not None
207-
else int(config["chunk_max_per_source_url"]),
208-
max_concurrent_crawls=request.max_concurrent_crawls
209-
or int(config["max_concurrent_crawls"]),
210-
max_concurrent_embedding_calls=request.max_concurrent_embedding_calls
211-
or int(config["max_concurrent_embedding_calls"]),
212-
embedding_timeout_seconds=request.embedding_timeout_seconds
213-
if request.embedding_timeout_seconds is not None
214-
else float(config["embedding_timeout_seconds"]),
215-
embedding_timeout_retries=request.embedding_timeout_retries
216-
if request.embedding_timeout_retries is not None
217-
else int(config["embedding_timeout_retries"]),
218-
crawl_max_chunk_tokens=request.crawl_max_chunk_tokens
219-
or int(config["crawl_max_chunk_tokens"]),
220-
crawl_overlap_tokens=request.crawl_overlap_tokens
221-
if request.crawl_overlap_tokens is not None
222-
else int(config["crawl_overlap_tokens"]),
223-
crawl_max_page_tokens=request.crawl_max_page_tokens
224-
if request.crawl_max_page_tokens is not None
225-
else int(config["crawl_max_page_tokens"]),
226-
crawl_fit_markdown_mode=(
227-
request.crawl_fit_markdown_mode
228-
if request.crawl_fit_markdown_mode is not None
229-
else str(config["crawl_fit_markdown_mode"])
230-
),
231-
crawl_fit_min_chars=(
232-
request.crawl_fit_min_chars
233-
if request.crawl_fit_min_chars is not None
234-
else int(config["crawl_fit_min_chars"])
235-
),
236-
crawl_bm25_threshold=request.crawl_bm25_threshold
237-
if request.crawl_bm25_threshold is not None
238-
else float(config["crawl_bm25_threshold"]),
239-
crawl_bm25_language=(
240-
request.crawl_bm25_language
241-
if request.crawl_bm25_language is not None
242-
else str(config["crawl_bm25_language"])
243-
),
244-
crawl_pruning_threshold=(
245-
request.crawl_pruning_threshold
246-
if request.crawl_pruning_threshold is not None
247-
else float(config["crawl_pruning_threshold"])
248-
),
249-
embedding_backend=str(config["embedding_backend"]),
250-
embedding_model=embedding_model,
251-
embedding_openai_env_file=str(config["embedding_openai_env_file"]),
252-
dense_query_prefix=request.dense_query_prefix
253-
if request.dense_query_prefix is not None
254-
else str(config["dense_query_prefix"]),
255-
dense_document_prefix=request.dense_document_prefix
256-
if request.dense_document_prefix is not None
257-
else str(config["dense_document_prefix"]),
258-
dense_document_embed_batch_size=request.dense_document_embed_batch_size
259-
if request.dense_document_embed_batch_size is not None
260-
else int(config["dense_document_embed_batch_size"]),
261-
blocked_domains=config["blocked_domains"],
262-
encoding_name=request.encoding_name or str(config["encoding_name"]),
263-
trace_path=Path(request.trace_path) if request.trace_path else config_trace_path(config),
197+
query,
198+
**run_kwargs,
199+
trace_path=Path(trace_path) if trace_path else config_trace_path(config),
264200
)
265201
return {"answer": result.answer}
266202

servers/mcp_server.py

Lines changed: 2 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
from services.research_config_service import (
2121
config_trace_path,
2222
load_research_config,
23+
normalize_research_query,
2324
research_run_kwargs,
2425
research_tokenizer_name,
2526
)
@@ -156,13 +157,6 @@ def _ensure_local_bundle_for_config(config: dict[str, Any]) -> None:
156157
ensure_onnx_bundle_sync(str(config["embedding_model"]))
157158

158159

159-
def _validate_query(query: str) -> str:
160-
query = query.strip()
161-
if not query:
162-
raise ValueError("query must not be empty")
163-
return query
164-
165-
166160
def _log(message: str) -> None:
167161
print(f"[tinysearch] {message}", file=sys.stderr, flush=True)
168162

@@ -198,7 +192,7 @@ def _enable_traceback_dump() -> None:
198192
),
199193
)
200194
async def research(query: str) -> dict[str, Any]:
201-
query = _validate_query(query)
195+
query = normalize_research_query(query)
202196
started = time.monotonic()
203197
_log(f"research called query={query!r}")
204198
try:

services/research_config_service.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -218,6 +218,13 @@ def research_run_kwargs(config: dict[str, Any] | None = None) -> dict[str, Any]:
218218
return {key: config[key] for key in keys}
219219

220220

221+
def normalize_research_query(query: str) -> str:
222+
query = query.strip()
223+
if not query:
224+
raise ValueError("query must not be empty")
225+
return query
226+
227+
221228
def research_embedding_model_info(config: dict[str, Any] | None = None) -> dict[str, str]:
222229
config = load_research_config() if config is None else config
223230
backend = normalize_embedding_backend(str(config["embedding_backend"]))

tests/test_server_embedding_startup.py

Lines changed: 77 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,22 @@
11
from __future__ import annotations
22

33
import unittest
4-
from unittest.mock import patch
5-
6-
from servers.fastapi_server import _ensure_local_bundle_for_config as ensure_fastapi_bundle
4+
from types import SimpleNamespace
5+
from unittest.mock import AsyncMock, patch
6+
7+
from servers.fastapi_server import (
8+
ResearchRequest,
9+
_ensure_local_bundle_for_config as ensure_fastapi_bundle,
10+
_tinysearch_version,
11+
research_endpoint,
12+
)
713
from servers.mcp_server import _ensure_local_bundle_for_config as ensure_mcp_bundle
14+
from services.research_config_service import (
15+
DEFAULT_RESEARCH_CONFIG,
16+
config_trace_path,
17+
normalize_research_query,
18+
research_run_kwargs,
19+
)
820

921

1022
class ServerEmbeddingStartupTests(unittest.IsolatedAsyncioTestCase):
@@ -25,6 +37,68 @@ async def test_fastapi_startup_skips_openai_compatible_backend(self) -> None:
2537
ensure.assert_not_called()
2638

2739

40+
class FastApiResearchParityTests(unittest.IsolatedAsyncioTestCase):
41+
async def test_research_uses_same_config_defaults_as_mcp(self) -> None:
42+
config = dict(DEFAULT_RESEARCH_CONFIG)
43+
config["embedding_backend"] = "openai_compatible"
44+
run = AsyncMock(return_value=SimpleNamespace(answer="grounded prompt"))
45+
46+
with patch(
47+
"servers.fastapi_server.load_research_config", return_value=config
48+
), patch("servers.fastapi_server.agentic_run", new=run):
49+
response = await research_endpoint(ResearchRequest(query=" test query "))
50+
51+
self.assertEqual(response, {"answer": "grounded prompt"})
52+
run.assert_awaited_once_with(
53+
"test query",
54+
**research_run_kwargs(config),
55+
trace_path=config_trace_path(config),
56+
)
57+
58+
async def test_research_request_overrides_shared_defaults(self) -> None:
59+
config = dict(DEFAULT_RESEARCH_CONFIG)
60+
config["embedding_backend"] = "openai_compatible"
61+
run = AsyncMock(return_value=SimpleNamespace(answer="grounded prompt"))
62+
63+
with patch(
64+
"servers.fastapi_server.load_research_config", return_value=config
65+
), patch("servers.fastapi_server.agentic_run", new=run):
66+
await research_endpoint(
67+
ResearchRequest(
68+
query="test query",
69+
search_top_k=7,
70+
pipeline_timeout_seconds=9.5,
71+
)
72+
)
73+
74+
kwargs = run.await_args.kwargs
75+
self.assertEqual(kwargs["search_top_k"], 7)
76+
self.assertEqual(kwargs["pipeline_timeout_seconds"], 9.5)
77+
78+
async def test_research_rejects_whitespace_only_query(self) -> None:
79+
with patch(
80+
"servers.fastapi_server.load_research_config",
81+
return_value=dict(DEFAULT_RESEARCH_CONFIG),
82+
):
83+
with self.assertRaisesRegex(ValueError, "query must not be empty"):
84+
await research_endpoint(ResearchRequest(query=" "))
85+
86+
87+
class ServerRuntimeMetadataTests(unittest.TestCase):
88+
def test_version_comes_from_environment(self) -> None:
89+
with patch.dict("os.environ", {"TINYSEARCH_VERSION": "v0.2.0"}):
90+
self.assertEqual(_tinysearch_version(), "v0.2.0")
91+
92+
def test_version_defaults_to_dev(self) -> None:
93+
with patch.dict("os.environ", {}, clear=True):
94+
self.assertEqual(_tinysearch_version(), "dev")
95+
96+
def test_query_normalization_is_shared(self) -> None:
97+
self.assertEqual(normalize_research_query(" hello "), "hello")
98+
with self.assertRaisesRegex(ValueError, "query must not be empty"):
99+
normalize_research_query(" ")
100+
101+
28102
class McpEmbeddingStartupTests(unittest.TestCase):
29103
def test_mcp_startup_ensures_selected_local_embedding_model(self) -> None:
30104
cfg = {"embedding_backend": "onnx", "embedding_model": "quality"}

0 commit comments

Comments
 (0)