SWE-Race › Tasks › joopsnijder-multi-search-api-2 ← prevnext →

joopsnijder-multi-search-api-2

Joopsnijder/multi-search-apicleansinglemerged 2025-11-03MITfix: 3 files, +95 −692 fail-to-pass · 36 pass-to-pass
Results
Modelsolved / attemptsmedian stepsmedian costattempts
GPT-5.6 Luna6/610$0.0091✓ 2✓ 3✓ 4✓ 5✓ 6✓
DeepSeek V4 Flash2/214$0.0051✓ 2✓
GLM-5.3 Flash2/26$0.0011✓ 2✓
The prompt the agent sees

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_concurrent_cache_reads_and_writestest_concurrent_cache_writes
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]

What each model did

Submitted diff first, then every command. Outputs show 12 lines; expand for the rest.
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]
step 1
$ 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
step 2
$ 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
step 3
$ cd /repo && python - <<'PY'
output · 1 lines
OK
step 4
$ 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(-)
step 5
$ 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 ==============================
step 6
$ 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.