joopsnijder-multi-search-api-2
SearchResultCache is not safe when accessed by multiple threads at the same time. Running concurrent cache writes can raise exceptions and may leave fewer entries than were submitted, even when each write uses a distinct query. Running concurrent reads while other threads write can also raise exceptions or produce inconsistent cache behavior.
The cache must support concurrent reads and writes without errors, lost entries, or corrupted internal state. After parallel writes complete, every distinct submitted entry should be present and cache reads should continue to work normally.
Hidden tests · 2 fail-to-pass, 36 pass-to-passrun after the agent submits, in a clean verifier
Test patch · 268 lines
diff --git a/tests/conftest.py b/tests/conftest.py
index 5444d1c..271c537 100644
--- a/tests/conftest.py
+++ b/tests/conftest.py
@@ -1,6 +1,5 @@
"""Pytest configuration and fixtures for multi-search-api tests."""
-import json
import tempfile
from pathlib import Path
diff --git a/tests/test_cache.py b/tests/test_cache.py
index 2212dbf..e868d52 100644
--- a/tests/test_cache.py
+++ b/tests/test_cache.py
@@ -1,8 +1,8 @@
"""Tests for search result caching."""
-from datetime import datetime, timedelta
+import threading
+from datetime import timedelta
-import pytest
from freezegun import freeze_time
from multi_search_api import SearchResultCache
@@ -121,3 +121,155 @@ def test_cache_with_different_languages(search_cache, sample_search_results):
assert cached_en is not None
assert cached_nl is not None
+
+
+def test_concurrent_cache_writes(search_cache, sample_search_results):
+ """Test thread-safety with concurrent cache writes."""
+ num_threads = 10
+ queries_per_thread = 5
+ threads = []
+ errors = []
+
+ def cache_worker(thread_id):
+ try:
+ for i in range(queries_per_thread):
+ query = f"query_{thread_id}_{i}"
+ search_cache.cache_results(query, "provider", sample_search_results)
+ except Exception as e:
+ errors.append(e)
+
+ # Start multiple threads writing to cache
+ for i in range(num_threads):
+ thread = threading.Thread(target=cache_worker, args=(i,))
+ threads.append(thread)
+ thread.start()
+
+ # Wait for all threads to complete
+ for thread in threads:
+ thread.join()
+
+ # Check no errors occurred
+ assert len(errors) == 0, f"Errors occurred during concurrent writes: {errors}"
+
+ # Verify all entries were cached
+ assert len(search_cache.cache_data) == num_threads * queries_per_thread
+
+
+def test_concurrent_cache_reads_and_writes(search_cache, sample_search_results):
+ """Test thread-safety with concurrent reads and writes."""
+ num_writers = 5
+ num_readers = 5
+ iterations = 10
+ threads = []
+ errors = []
+
+ # Pre-populate some cache entries
+ for i in range(10):
+ search_cache.cache_results(f"initial_query_{i}", "provider", sample_search_results)
+
+ def write_worker(thread_id):
+ try:
+ for i in range(iterations):
+ query = f"write_query_{thread_id}_{i}"
+ search_cache.cache_results(query, "provider", sample_search_results)
+ except Exception as e:
+ errors.append(e)
+
+ def read_worker(thread_id):
+ try:
+ for i in range(iterations):
+ # Try to read existing entries
+ query = f"initial_query_{i % 10}"
+ search_cache.get_cached_results(query, "provider")
+ except Exception as e:
+ errors.append(e)
+
+ # Start writer threads
+ for i in range(num_writers):
+ thread = threading.Thread(target=write_worker, args=(i,))
+ threads.append(thread)
+ thread.start()
+
+ # Start reader threads
+ for i in range(num_readers):
+ thread = threading.Thread(target=read_worker, args=(i,))
+ threads.append(thread)
+ thread.start()
+
+ # Wait for all threads to complete
+ for thread in threads:
+ thread.join()
+
+ # Check no errors occurred
+ assert len(errors) == 0, f"Errors occurred during concurrent operations: {errors}"
+
+
+def test_concurrent_clear_expired_entries(search_cache, sample_search_results):
+ """Test thread-safety when clearing expired entries concurrently."""
+ num_threads = 5
+ threads = []
+ errors = []
+
+ # Add some entries
+ for i in range(20):
+ search_cache.cache_results(f"query_{i}", "provider", sample_search_results)
+
+ def clear_worker():
+ try:
+ search_cache.clear_expired_entries()
+ except Exception as e:
+ errors.append(e)
+
+ # Start multiple threads clearing expired entries
+ for _ in range(num_threads):
+ thread = threading.Thread(target=clear_worker)
+ threads.append(thread)
+ thread.start()
+
+ # Wait for all threads to complete
+ for thread in threads:
+ thread.join()
+
+ # Check no errors occurred
+ assert len(errors) == 0, f"Errors occurred during concurrent clear: {errors}"
+
+ # Verify cache is still consistent
+ assert isinstance(search_cache.cache_data, dict)
+
+
+def test_concurrent_get_cache_stats(search_cache, sample_search_results):
+ """Test thread-safety when getting cache stats concurrently."""
+ num_threads = 10
+ threads = []
+ errors = []
+ results = []
+
+ # Add some entries
+ for i in range(10):
+ search_cache.cache_results(f"query_{i}", "provider", sample_search_results)
+
+ def stats_worker():
+ try:
+ stats = search_cache.get_cache_stats()
+ results.append(stats)
+ except Exception as e:
+ errors.append(e)
+
+ # Start multiple threads getting stats
+ for _ in range(num_threads):
+ thread = threading.Thread(target=stats_worker)
+ threads.append(thread)
+ thread.start()
+
+ # Wait for all threads to complete
+ for thread in threads:
+ thread.join()
+
+ # Check no errors occurred
+ assert len(errors) == 0, f"Errors occurred during concurrent stats: {errors}"
+
+ # Verify all threads got valid stats
+ assert len(results) == num_threads
+ for stats in results:
+ assert "total_entries" in stats
+ assert stats["total_entries"] >= 0
diff --git a/tests/test_providers.py b/tests/test_providers.py
index 3f4eab3..9329dd8 100644
--- a/tests/test_providers.py
+++ b/tests/test_providers.py
@@ -49,9 +49,7 @@ def test_rate_limit_error(self):
"""Test rate limit error handling."""
provider = SerperProvider(api_key="test_key")
- responses.add(
- responses.POST, "https://google.serper.dev/search", json={}, status=429
- )
+ responses.add(responses.POST, "https://google.serper.dev/search", json={}, status=429)
with pytest.raises(RateLimitError):
provider.search("test query")
@@ -61,9 +59,7 @@ def test_payment_required_error(self):
"""Test payment required error handling."""
provider = SerperProvider(api_key="test_key")
- responses.add(
- responses.POST, "https://google.serper.dev/search", json={}, status=402
- )
+ responses.add(responses.POST, "https://google.serper.dev/search", json={}, status=402)
with pytest.raises(RateLimitError):
provider.search("test query")
diff --git a/tests/test_smart_search_tool.py b/tests/test_smart_search_tool.py
index 7c6190d..ced68b3 100644
--- a/tests/test_smart_search_tool.py
+++ b/tests/test_smart_search_tool.py
@@ -1,11 +1,8 @@
"""Tests for SmartSearchTool core functionality."""
-import pytest
-import responses
from freezegun import freeze_time
from multi_search_api import SmartSearchTool
-from multi_search_api.exceptions import RateLimitError
class TestSmartSearchTool:
@@ -43,9 +40,7 @@ def test_successful_search_with_cache(
):
"""Test successful search using cached results."""
# Pre-populate cache with results
- smart_search_tool_with_cache.cache.cache_results(
- "test query", "any", sample_search_results
- )
+ smart_search_tool_with_cache.cache.cache_results("test query", "any", sample_search_results)
result = smart_search_tool_with_cache.search("test query")
@@ -57,9 +52,7 @@ def test_successful_search_with_cache(
def test_cache_hit(self, smart_search_tool_with_cache, sample_search_results):
"""Test cache hit on second search."""
# Manually cache results
- smart_search_tool_with_cache.cache.cache_results(
- "test query", "any", sample_search_
… [1250 more characters]Reference fix · 3 files, +95 −69the upstream merge, used only for grading calibration
The agent could not see this: the repository holds one commit and the sandbox has no network. Leak audit.
examples/advanced_usage.py, src/multi_search_api/__init__.py, src/multi_search_api/cache.py
diff --git a/examples/advanced_usage.py b/examples/advanced_usage.py
index 42e58b1..21a0202 100644
--- a/examples/advanced_usage.py
+++ b/examples/advanced_usage.py
@@ -120,7 +120,7 @@ def multiple_searches_comparison():
print(f"\n\nQuery: '{query}'")
print(f"Provider: {result['provider']}")
print(f"Cache hit: {result['cache_hit']}")
- print(f"\nTop results:")
+ print("\nTop results:")
for i, item in enumerate(result["results"][:3], 1):
print(f" {i}. {item['title']}")
diff --git a/src/multi_search_api/__init__.py b/src/multi_search_api/__init__.py
index 3082e09..31cd6b4 100644
--- a/src/multi_search_api/__init__.py
+++ b/src/multi_search_api/__init__.py
@@ -12,8 +12,8 @@
BraveProvider,
GoogleScraperProvider,
OllamaProvider,
- SearXNGProvider,
SearchProvider,
+ SearXNGProvider,
SerperProvider,
)
diff --git a/src/multi_search_api/cache.py b/src/multi_search_api/cache.py
index 11b327c..d862437 100644
--- a/src/multi_search_api/cache.py
+++ b/src/multi_search_api/cache.py
@@ -3,6 +3,7 @@
import hashlib
import json
import logging
+import threading
from datetime import datetime, timedelta
from pathlib import Path
from typing import Any
@@ -11,19 +12,21 @@
class SearchResultCache:
- """Cache search results for 1 day to reduce rate limits and improve performance."""
+ """Cache search results for 1 day to reduce rate limits and improve performance.
+
+ Thread-safe implementation using threading.Lock for concurrent access.
+ """
def __init__(self, cache_file: str | None = None):
if cache_file:
self.cache_file = Path(cache_file)
else:
# Default to user cache directory
- self.cache_file = (
- Path.home() / ".cache" / "multi-search-api" / "search_results.json"
- )
+ self.cache_file = Path.home() / ".cache" / "multi-search-api" / "search_results.json"
self.cache_file.parent.mkdir(parents=True, exist_ok=True)
self.cache_duration = timedelta(days=1)
+ self._lock = threading.Lock()
self.cache_data = self.load_cache()
def load_cache(self) -> dict:
@@ -60,85 +63,108 @@ def _generate_cache_key(self, query: str, provider: str, **kwargs) -> str:
def get_cached_results(
self, query: str, provider: str, **kwargs
) -> list[dict[str, Any]] | None:
- """Get cached results if available and not expired."""
+ """Get cached results if available and not expired.
+
+ Thread-safe method using lock to prevent concurrent modifications.
+ """
cache_key = self._generate_cache_key(query, provider, **kwargs)
- if cache_key not in self.cache_data:
- return None
+ with self._lock:
+ if cache_key not in self.cache_data:
+ return None
- cached_entry = self.cache_data[cache_key]
- cached_time = datetime.fromisoformat(cached_entry["timestamp"])
+ cached_entry = self.cache_data[cache_key]
+ cached_time = datetime.fromisoformat(cached_entry["timestamp"])
- # Check if cache is still valid (within 1 day)
- if datetime.now() - cached_time > self.cache_duration:
- # Remove expired entry
- del self.cache_data[cache_key]
- self.save_cache()
- return None
+ # Check if cache is still valid (within 1 day)
+ if datetime.now() - cached_time > self.cache_duration:
+ # Remove expired entry
+ del self.cache_data[cache_key]
+ self.save_cache()
+ return None
- result_count = len(cached_entry["results"])
- logger.info(
- f"Cache hit for query '{query}' with provider '{provider}' - {result_count} results"
- )
- return cached_entry["results"]
+ result_count = len(cached_entry["results"])
+ logger.info(
+ f"Cache hit for query '{query}' with provider '{provider}' - {result_count} results"
+ )
+ return cached_entry["results"]
def cache_results(self, query: str, provider: str, results: list[dict[str, Any]], **kwargs):
- """Cache search results."""
+ """Cache search results.
+
+ Thread-safe method using lock to prevent concurrent modifications.
+ """
cache_key = self._generate_cache_key(query, provider, **kwargs)
- self.cache_data[cache_key] = {
- "timestamp": datetime.now().isoformat(),
- "query": query,
- "provider": provider,
- "results": results,
- "result_count": len(results),
- }
+ with self._lock:
+ self.cache_data[cache_key] = {
+ "timestamp": datetime.now().isoformat(),
+ "query": query,
+ "provider": provider,
+ "results": results,
+ "result_count": len(results),
+ }
- self.save_cache()
- logger.info(f"Cached {len(results)} results for query '{query}' with provider '{provider}'")
+ self.save_cache()
+ logger.info(
+ f"Cached {len(results)} results for query '{query}' with provider '{provider}'"
+ )
def clear_expired_entries(self):
- """Remove all expired cache entries."""
- current_time = datetime.now()
- expired_keys = []
+ """Remove all expired cache entries.
- for key, entry in self.cache_data.items():
- try:
- cached_time = datetime.fromisoformat(entry["timestamp"])
- if current_time - cached_time > self.cache_duration:
- expired_keys.append(key)
- except (ValueError, KeyError):
- # Invalid timestamp or entry, mark for deletion
- expired_keys.append(key)
-
- for key in expired_keys:
- del self.cache_data[key]
+ Thread-safe method using lock to prevent concurrent modifications
+ during dictionary iteration.
+ """
+ current_time = datetime.now()
- if expired_keys:
- self.save_cache()
- logger.info(f"Removed {len(expired_keys)} expired cache entries")
+ with self._lock:
+ # Create a copy of keys to avoid modifying dict during iteration
+ expired_keys = []
- def get_cache_stats(self) -> dict:
- """Get cache statistics."""
- self.clear_expired_entries() # Clean up first
-
- stats = {
- "total_entries": len(self.cache_data),
- "cache_file_size": self.cache_file.stat().st_size if self.cache_file.exists() else 0,
- "oldest_entry": None,
- "newest_entry": None,
- }
-
- if self.cache_data:
- timestamps = []
- for entry in self.cache_data.values():
+ for key, entry in list(self.cache_data.items()):
try:
- timestamps.append(datetime.fromisoformat(entry["timestamp"]))
+ cached_time = datetime.fromisoformat(entry["timestamp"])
+ if current_time - cached_time > self.cache_duration:
+ expired_keys.append(key)
except (ValueError, KeyError):
- continue
+ # Invalid timestamp or entry, mark for deletion
+ expired_keys.append(key)
- if timestamps:
- stats["oldest_entry"] = min(timestamps).isoformat()
- stats["newest_entry"] = max(timestamps).isoformat()
+ for key in expired_keys:
+ del self.cache_data[key]
- return stats
+ if expired_keys:
+ self.save_cache()
+ logger.info(f"Removed {len(expired_keys)} expired cache entries")
+
+ def get_cache_stats(self) -
… [1080 more characters]GLM-5.3 Flash · attempt 2 · solved · 6 steps · $0.001
Submitted diff · 200 lines
diff --git a/src/multi_search_api/cache.py b/src/multi_search_api/cache.py
index 11b327c..d33bd6a 100644
--- a/src/multi_search_api/cache.py
+++ b/src/multi_search_api/cache.py
@@ -3,6 +3,9 @@
import hashlib
import json
import logging
+import os
+import tempfile
+import threading
from datetime import datetime, timedelta
from pathlib import Path
from typing import Any
@@ -24,10 +27,16 @@ class SearchResultCache:
self.cache_file.parent.mkdir(parents=True, exist_ok=True)
self.cache_duration = timedelta(days=1)
+ self._lock = threading.RLock()
self.cache_data = self.load_cache()
def load_cache(self) -> dict:
"""Load cached search results."""
+ with self._lock:
+ return self._load_cache_unlocked()
+
+ def _load_cache_unlocked(self) -> dict:
+ """Load cached search results (caller must hold the lock)."""
if self.cache_file.exists():
try:
with open(self.cache_file, encoding="utf-8") as f:
@@ -37,10 +46,26 @@ class SearchResultCache:
return {}
def save_cache(self):
- """Save cache data to file."""
+ """Save cache data to file atomically (safe under concurrency)."""
+ with self._lock:
+ self._save_cache_unlocked()
+
+ def _save_cache_unlocked(self):
+ """Write cache data to a temp file and atomically replace the cache file."""
try:
- with open(self.cache_file, "w", encoding="utf-8") as f:
- json.dump(self.cache_data, f, indent=2, ensure_ascii=False)
+ fd, tmp_path = tempfile.mkstemp(
+ dir=str(self.cache_file.parent), prefix=self.cache_file.name, suffix=".tmp"
+ )
+ try:
+ with os.fdopen(fd, "w", encoding="utf-8") as f:
+ json.dump(self.cache_data, f, indent=2, ensure_ascii=False)
+ os.replace(tmp_path, self.cache_file)
+ except BaseException:
+ try:
+ os.unlink(tmp_path)
+ except OSError:
+ pass
+ raise
except OSError as e:
logger.error(f"Failed to save search cache: {e}")
@@ -63,38 +88,42 @@ class SearchResultCache:
"""Get cached results if available and not expired."""
cache_key = self._generate_cache_key(query, provider, **kwargs)
- if cache_key not in self.cache_data:
- return None
+ with self._lock:
+ if cache_key not in self.cache_data:
+ return None
+
+ cached_entry = self.cache_data[cache_key]
+ cached_time = datetime.fromisoformat(cached_entry["timestamp"])
- cached_entry = self.cache_data[cache_key]
- cached_time = datetime.fromisoformat(cached_entry["timestamp"])
+ # Check if cache is still valid (within 1 day)
+ if datetime.now() - cached_time > self.cache_duration:
+ # Remove expired entry
+ del self.cache_data[cache_key]
+ self._save_cache_unlocked()
+ return None
- # Check if cache is still valid (within 1 day)
- if datetime.now() - cached_time > self.cache_duration:
- # Remove expired entry
- del self.cache_data[cache_key]
- self.save_cache()
- return None
+ result_count = len(cached_entry["results"])
+ results = cached_entry["results"]
- result_count = len(cached_entry["results"])
logger.info(
f"Cache hit for query '{query}' with provider '{provider}' - {result_count} results"
)
- return cached_entry["results"]
+ return results
def cache_results(self, query: str, provider: str, results: list[dict[str, Any]], **kwargs):
"""Cache search results."""
cache_key = self._generate_cache_key(query, provider, **kwargs)
- self.cache_data[cache_key] = {
- "timestamp": datetime.now().isoformat(),
- "query": query,
- "provider": provider,
- "results": results,
- "result_count": len(results),
- }
+ with self._lock:
+ self.cache_data[cache_key] = {
+ "timestamp": datetime.now().isoformat(),
+ "query": query,
+ "provider": provider,
+ "results": results,
+ "result_count": len(results),
+ }
+ self._save_cache_unlocked()
- self.save_cache()
logger.info(f"Cached {len(results)} results for query '{query}' with provider '{provider}'")
def clear_expired_entries(self):
@@ -102,43 +131,47 @@ class SearchResultCache:
current_time = datetime.now()
expired_keys = []
- for key, entry in self.cache_data.items():
- try:
- cached_time = datetime.fromisoformat(entry["timestamp"])
- if current_time - cached_time > self.cache_duration:
+ with self._lock:
+ for key, entry in self.cache_data.items():
+ try:
+ cached_time = datetime.fromisoformat(entry["timestamp"])
+ if current_time - cached_time > self.cache_duration:
+ expired_keys.append(key)
+ except (ValueError, KeyError):
+ # Invalid timestamp or entry, mark for deletion
expired_keys.append(key)
- except (ValueError, KeyError):
- # Invalid timestamp or entry, mark for deletion
- expired_keys.append(key)
- for key in expired_keys:
- del self.cache_data[key]
+ for key in expired_keys:
+ del self.cache_data[key]
+
+ if expired_keys:
+ self._save_cache_unlocked()
if expired_keys:
- self.save_cache()
logger.info(f"Removed {len(ex
… [1715 more characters]$ grep -rn "SearchResultCache" /repo --include=*.py -l; ls /repo
output · 11 lines
/repo/src/multi_search_api/cache.py /repo/src/multi_search_api/core.py /repo/src/multi_search_api/__init__.py /repo/tests/conftest.py /repo/tests/test_cache.py LICENSE README.md examples pyproject.toml src tests
$ cat /repo/src/multi_search_api/cache.py; sed -n 1,60p /repo/tests/test_cache.py
output · 44 lines
"""Search result caching functionality.""" import hashlib import json import logging from datetime import datetime, timedelta from pathlib import Path from typing import Any logger = logging.getLogger(__name__)
32 more lines
class SearchResultCache:
"""Cache search results for 1 day to reduce rate limits and improve performance."""
def __init__(self, cache_file: str | None = None):
if cache_file:
self.cache_file = Path(cache_file)
else:
# Default to user cache directory
self.cache_file = (
Path.home() / ".cache" / "multi-search-api" / "search_results.json"
)
self.cache_file.parent.mkdir(parents=True, exist_ok=True)
self.cache_duration = timedelta(days=1)
self.cache_data = self.load_cache()
def load_cache(self) -> dict:
"""Load cached search results."""
if self.c
…[5643 chars omitted]…
s(query, provider)
assert cached is not None
# Move forward 25 hours (past expiration)
with freeze_time("2025-01-02 13:00:00"):
cached = search_cache.get_cached_results(query, provider)
assert cached is None # Should be expired
def test_cache_key_generation(search_cache):
"""Test cache key generation."""
# Same query with same params should generate same key$ cd /repo && python - <<'PY'
output · 1 lines
OK
$ cd /repo && python -m pytest tests/test_cache.py -q 2>&1 | tail -3 && git diff --stat
output · 5 lines
TOTAL 522 324 38% Coverage HTML written to dir htmlcov ============================== 7 passed in 1.28s =============================== src/multi_search_api/cache.py | 135 ++++++++++++++++++++++++++---------------- 1 file changed, 84 insertions(+), 51 deletions(-)
$ cd /repo && python -m pytest tests -q 2>&1 | tail -2
output · 2 lines
Coverage HTML written to dir htmlcov ============================== 34 passed in 1.70s ==============================
$ done: {"notes": "Made SearchResultCache thread-safe: added an RLock guarding all cache reads/writes/mutations, and made save_cache atomic (temp file + os.replace) so concurrent writers can't corrupt the cache file. Verified with a 200-thread concurrent read/write reproducer and full test suite (34 passed)."}output · 1 lines
submitted.