behavior-in-the-wild's picture
Deploy SDR-Arena leaderboard
f9e2361 verified
Raw
History Blame Contribute Delete
81.7 kB
"""
Web Search Service API
Centralized web search proxy with provider abstraction, date filtering,
and SQLite-based caching.
"""
import asyncio
import json
import hashlib
import uuid
from datetime import datetime, timedelta
from typing import Optional, List, Dict, Any
from contextlib import asynccontextmanager
from fastapi import FastAPI, HTTPException, Query, Request
from fastapi.middleware.cors import CORSMiddleware
from starlette.middleware.base import BaseHTTPMiddleware
from pydantic import BaseModel, Field
from config import config, WebSearchConfig
from providers.base import SearchResult, QueueFullError
from tracing import get_request_id, set_request_id, reset_request_id, log
# =============================================================================
# Request Tracing Middleware
# =============================================================================
class RequestIDMiddleware(BaseHTTPMiddleware):
"""
Middleware to add unique request IDs to each incoming request.
- Generates a UUID for each request
- Stores it in context variable for use throughout request lifecycle
- Adds X-Request-ID header to response
"""
async def dispatch(self, request: Request, call_next):
# Generate unique request ID (or use incoming X-Request-ID if provided)
request_id = request.headers.get('X-Request-ID', str(uuid.uuid4()))
# Set the context variable
token = set_request_id(request_id)
try:
# Log request start (only for meaningful endpoints)
if request.url.path not in ['/health', '/favicon.ico']:
log(f"→ {request.method} {request.url.path}")
# Process request
response = await call_next(request)
# Add request ID to response headers
response.headers['X-Request-ID'] = request_id
# Log request end (only for meaningful endpoints)
if request.url.path not in ['/health', '/favicon.ico']:
log(f"← {request.method} {request.url.path} [{response.status_code}]")
return response
finally:
# Reset context variable
reset_request_id(token)
# =============================================================================
# Pydantic Models
# =============================================================================
class SearchRequest(BaseModel):
"""Request model for search endpoint."""
queries: List[str] = Field(..., description="One or more search queries")
provider: Optional[str] = Field(None, description="Search provider (tavily, google, brave)")
max_results: int = Field(default=5, ge=1, le=20, description="Max results per query")
start_date: Optional[str] = Field(None, description="Start date filter (YYYY-MM-DD)")
end_date: Optional[str] = Field(None, description="End date filter (YYYY-MM-DD)")
use_cache: bool = Field(default=True, description="Enable/disable caching")
filter_mode: str = Field(
default="none",
description="Post-processing filter mode: 'none' (default), 'heuristic' (URL-based), 'llm' (AI-powered)"
)
class SearchResultResponse(BaseModel):
"""Individual search result."""
url: str
title: str
content: str
snippet: str
raw_content: Optional[str] = None
score: Optional[float] = None
published_date: Optional[str] = None
filtered: bool = Field(default=False, description="Whether result was filtered out")
filter_reason: Optional[str] = Field(default=None, description="Reason for filtering")
# Content fetch details
fetch_method: Optional[str] = Field(default=None, description="How content was fetched: simple, crawl4ai, proxy, skip, none")
fetch_time_ms: Optional[float] = Field(default=None, description="Time taken to fetch content in ms")
fetch_error: Optional[str] = Field(default=None, description="Error message if content fetch failed")
class FilterStats(BaseModel):
"""Statistics about filtering applied."""
mode: str
total_before: int
total_after: int
filtered_count: int
filtered_urls: List[str] = []
class SearchResponse(BaseModel):
"""Response model for search endpoint."""
query: str
provider: str
results: List[SearchResultResponse]
cached: bool
search_time_ms: float # Query/search time only
content_fetch_time_ms: float = Field(default=0.0, description="Time spent fetching content")
total_time_ms: float = Field(default=0.0, description="Total time (search + content fetch)")
class BatchSearchResponse(BaseModel):
"""Response for batch search requests."""
searches: List[SearchResponse]
total_time_ms: float
cache_hits: int
cache_misses: int
filter_stats: Optional[FilterStats] = None
class ProviderInfo(BaseModel):
"""Information about a search provider."""
name: str
available: bool
supports_date_filter: bool
requires_api_key: bool
class HealthResponse(BaseModel):
"""Health check response."""
status: str
service: str
version: str
providers: List[ProviderInfo]
cache_enabled: bool
default_provider: str
class CacheStatsResponse(BaseModel):
"""Cache statistics response."""
total_entries: int
cache_size_bytes: int
oldest_entry: Optional[str]
newest_entry: Optional[str]
hit_count: int
# =============================================================================
# Date Filtering
# =============================================================================
import re
import httpx
class DateFilter:
"""
Post-processing date filter for search results.
Supports two modes:
- heuristic: Fast URL-based filtering (checks for year patterns in URLs)
- llm: AI-powered content analysis (more accurate but slower)
"""
def __init__(self, start_date: Optional[str], end_date: Optional[str]):
self.start_date = start_date
self.end_date = end_date
self.start_year = int(start_date[:4]) if start_date else None
self.end_year = int(end_date[:4]) if end_date else None
# Pattern to match years in URLs
self.year_pattern = re.compile(r'/(\d{4})/')
self.year_in_url_pattern = re.compile(r'20\d{2}')
def _extract_year_from_url(self, url: str) -> Optional[int]:
"""Extract a year from URL path if present."""
# Try to find year in URL path (e.g., /2024/01/article)
match = self.year_pattern.search(url)
if match:
year = int(match.group(1))
if 1990 <= year <= 2100: # Reasonable year range
return year
# Also check for date patterns like 2024-01-15 in URL
date_match = re.search(r'(\d{4})-\d{2}-\d{2}', url)
if date_match:
year = int(date_match.group(1))
if 1990 <= year <= 2100:
return year
return None
def filter_heuristic(self, results: List[SearchResult]) -> tuple[List[SearchResult], List[dict]]:
"""
Apply heuristic URL-based filtering.
Filters out results where the URL contains a year outside the date range.
This catches obvious violations like /2026/01/article when filtering for 2020-2022.
Returns: (filtered_results, filter_info)
"""
if not self.start_year or not self.end_year:
return results, []
filtered = []
removed = []
for result in results:
url_year = self._extract_year_from_url(result.url)
if url_year is not None:
if url_year < self.start_year or url_year > self.end_year:
removed.append({
"url": result.url,
"reason": f"URL contains year {url_year} outside range {self.start_year}-{self.end_year}"
})
continue
filtered.append(result)
return filtered, removed
async def filter_llm(
self,
results: List[SearchResult],
vllm_url: Optional[str] = None,
model: str = "Qwen/Qwen2.5-72B-Instruct"
) -> tuple[List[SearchResult], List[dict]]:
"""
Apply LLM-based content filtering.
Uses an LLM to analyze each result's content and determine if it's
actually from within the date range (not just published then, but
discussing events from that time period).
Returns: (filtered_results, filter_info)
"""
if not self.start_date or not self.end_date:
return results, []
if not vllm_url:
# Fall back to heuristic if no LLM available
return self.filter_heuristic(results)
filtered = []
removed = []
# Process in batches to avoid overwhelming the LLM
async with httpx.AsyncClient(timeout=30.0) as client:
for result in results:
# Prepare content for analysis
content = result.content or result.snippet or ""
if len(content) > 1500:
content = content[:1500] + "..."
prompt = f"""Analyze if this search result content is appropriate for research constrained to the time period {self.start_date} to {self.end_date}.
URL: {result.url}
Title: {result.title}
Content: {content}
Consider:
1. Does the URL suggest a publication date outside the range?
2. Does the content primarily discuss events, data, or information from AFTER {self.end_date}?
3. Is this content about future predictions/projections written during the valid period (this is OK)?
Answer with exactly one word: KEEP or FILTER
If FILTER, add a brief reason on a new line."""
try:
response = await client.post(
f"{vllm_url}/v1/chat/completions",
json={
"model": model,
"messages": [{"role": "user", "content": prompt}],
"max_tokens": 50,
"temperature": 0.1
},
headers={"Authorization": "Bearer not-needed"}
)
if response.status_code == 200:
data = response.json()
answer = data["choices"][0]["message"]["content"].strip()
if answer.upper().startswith("FILTER"):
reason = answer.split("\n", 1)[1] if "\n" in answer else "LLM determined content outside date range"
removed.append({
"url": result.url,
"reason": reason
})
continue
filtered.append(result)
except Exception as e:
# On error, keep the result
print(f"LLM filter error for {result.url}: {e}")
filtered.append(result)
return filtered, removed
# =============================================================================
# Provider Registry
# =============================================================================
class ProviderRegistry:
"""Registry for search providers with lazy initialization."""
def __init__(self, cfg: WebSearchConfig):
self.config = cfg
self._providers: Dict[str, Any] = {}
self._initialized = False
async def initialize(self):
"""Initialize all available providers."""
if self._initialized:
return
# Import providers lazily to avoid import errors if dependencies missing
if self.config.tavily_api_key:
try:
from providers.tavily import TavilyProvider
self._providers["tavily"] = TavilyProvider(
api_key=self.config.tavily_api_key,
max_concurrent=self.config.max_concurrent,
max_queue_size=self.config.max_queue_size
)
except ImportError as e:
print(f"Warning: Could not load Tavily provider: {e}")
if self.config.brave_api_key:
try:
from providers.brave import BraveProvider
self._providers["brave"] = BraveProvider(
api_key=self.config.brave_api_key,
max_concurrent=self.config.max_concurrent,
max_queue_size=self.config.max_queue_size
)
except ImportError as e:
print(f"Warning: Could not load Brave provider: {e}")
if self.config.google_api_key and self.config.google_cx:
try:
from providers.google import GoogleProvider
self._providers["google"] = GoogleProvider(
api_key=self.config.google_api_key,
cx=self.config.google_cx,
max_concurrent=self.config.max_concurrent,
max_queue_size=self.config.max_queue_size
)
except ImportError as e:
print(f"Warning: Could not load Google provider: {e}")
if self.config.brightdata_api_token:
try:
from providers.brightdata import BrightdataProvider
self._providers["brightdata"] = BrightdataProvider(
api_token=self.config.brightdata_api_token,
zone=self.config.brightdata_serp_zone,
max_concurrent=self.config.max_concurrent, # Increased to 50 by default
max_queue_size=self.config.max_queue_size, # Backpressure queue limit
output_format="json", # Use JSON for structured results
proxy_user=self.config.brightdata_proxy_user,
proxy_password=self.config.brightdata_proxy_password,
enable_proxy_fallback=self.config.content_proxy_enabled, # Default: False for speed
proxy_timeout=self.config.content_proxy_timeout, # Default: 10s
parallel_timeout=self.config.content_parallel_timeout, # Default: 12s
)
except ImportError as e:
print(f"Warning: Could not load Brightdata provider: {e}")
self._initialized = True
def get_provider(self, name: Optional[str] = None):
"""Get a provider by name or return default."""
provider_name = name or self.config.default_provider
if provider_name not in self._providers:
available = list(self._providers.keys())
if not available:
raise HTTPException(
status_code=503,
detail="No search providers available. Check API key configuration."
)
# Fall back to first available
provider_name = available[0]
return self._providers[provider_name], provider_name
def get_all_providers_info(self) -> List[ProviderInfo]:
"""Get info about all providers."""
infos = []
for name, provider in self._providers.items():
info = provider.get_info()
infos.append(ProviderInfo(**info))
return infos
async def close_all(self):
"""Close all provider connections."""
for provider in self._providers.values():
await provider.close()
# =============================================================================
# Cache Manager
# =============================================================================
class CacheManager:
"""SQLite-based cache manager with date-aware keys and proper concurrency control."""
def __init__(self, db_path: str, ttl_hours: int = 24, enabled: bool = True):
self.db_path = db_path
self.ttl_hours = ttl_hours
self.enabled = enabled # Global enable/disable flag
self._conn = None
self._lock = asyncio.Lock() # Serialize all database operations
async def initialize(self):
"""Initialize the database connection and schema."""
if not self.enabled:
print("[CacheManager] Cache DISABLED - skipping initialization", flush=True)
return
import aiosqlite
import os
# Ensure directory exists
os.makedirs(os.path.dirname(self.db_path), exist_ok=True)
self._conn = await aiosqlite.connect(self.db_path)
# Enable WAL mode for better concurrency
await self._conn.execute("PRAGMA journal_mode=WAL")
# Busy timeout: wait up to 30 seconds before erroring on lock
await self._conn.execute("PRAGMA busy_timeout=30000")
# Synchronous mode: normal (balance between safety and speed)
await self._conn.execute("PRAGMA synchronous=NORMAL")
# Create table if not exists
await self._conn.execute("""
CREATE TABLE IF NOT EXISTS search_cache (
cache_key TEXT PRIMARY KEY,
query TEXT NOT NULL,
provider TEXT NOT NULL,
start_date TEXT,
end_date TEXT,
max_results INTEGER NOT NULL,
results TEXT NOT NULL,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
expires_at TIMESTAMP NOT NULL,
hit_count INTEGER DEFAULT 0
)
""")
# Create index for cleanup
await self._conn.execute("""
CREATE INDEX IF NOT EXISTS idx_expires_at ON search_cache(expires_at)
""")
await self._conn.commit()
def _make_cache_key(
self,
query: str,
provider: str,
start_date: Optional[str],
end_date: Optional[str],
max_results: int
) -> str:
"""Generate a unique cache key."""
key_string = f"{query}|{provider}|{start_date or 'any'}|{end_date or 'any'}|{max_results}"
return hashlib.sha256(key_string.encode()).hexdigest()
async def get(
self,
query: str,
provider: str,
start_date: Optional[str],
end_date: Optional[str],
max_results: int
) -> Optional[List[SearchResult]]:
"""Get cached results if available and not expired."""
if not self.enabled or not self._conn:
return None
cache_key = self._make_cache_key(query, provider, start_date, end_date, max_results)
async with self._lock:
try:
cursor = await self._conn.execute(
"""
SELECT results FROM search_cache
WHERE cache_key = ? AND expires_at > datetime('now')
""",
(cache_key,)
)
row = await cursor.fetchone()
if row:
# Update hit count
await self._conn.execute(
"UPDATE search_cache SET hit_count = hit_count + 1 WHERE cache_key = ?",
(cache_key,)
)
await self._conn.commit()
# Parse results
results_data = json.loads(row[0])
return [SearchResult.from_dict(r) for r in results_data]
return None
except Exception as e:
# Log but don't fail on cache errors
print(f"[CacheManager] GET error: {e}", flush=True)
return None
async def set(
self,
query: str,
provider: str,
start_date: Optional[str],
end_date: Optional[str],
max_results: int,
results: List[SearchResult]
):
"""Cache search results."""
if not self.enabled or not self._conn:
return
cache_key = self._make_cache_key(query, provider, start_date, end_date, max_results)
expires_at = datetime.utcnow() + timedelta(hours=self.ttl_hours)
results_json = json.dumps([r.to_dict() for r in results])
async with self._lock:
try:
await self._conn.execute(
"""
INSERT OR REPLACE INTO search_cache
(cache_key, query, provider, start_date, end_date, max_results, results, expires_at, hit_count)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0)
""",
(cache_key, query, provider, start_date, end_date, max_results, results_json, expires_at)
)
await self._conn.commit()
except Exception as e:
# Log but don't fail on cache errors
print(f"[CacheManager] SET error: {e}", flush=True)
async def get_stats(self) -> CacheStatsResponse:
"""Get cache statistics."""
if not self.enabled or not self._conn:
return CacheStatsResponse(
total_entries=0,
cache_size_bytes=0,
oldest_entry=None,
newest_entry=None,
hit_count=0
)
async with self._lock:
try:
# Total entries
cursor = await self._conn.execute("SELECT COUNT(*) FROM search_cache")
total_entries = (await cursor.fetchone())[0]
# Cache size (approximate)
cursor = await self._conn.execute(
"SELECT SUM(LENGTH(results)) FROM search_cache"
)
row = await cursor.fetchone()
cache_size = row[0] or 0
# Oldest and newest
cursor = await self._conn.execute(
"SELECT MIN(created_at), MAX(created_at) FROM search_cache"
)
row = await cursor.fetchone()
oldest = row[0]
newest = row[1]
# Total hits
cursor = await self._conn.execute("SELECT SUM(hit_count) FROM search_cache")
row = await cursor.fetchone()
hit_count = row[0] or 0
return CacheStatsResponse(
total_entries=total_entries,
cache_size_bytes=cache_size,
oldest_entry=oldest,
newest_entry=newest,
hit_count=hit_count
)
except Exception as e:
print(f"[CacheManager] STATS error: {e}", flush=True)
return CacheStatsResponse(
total_entries=0,
cache_size_bytes=0,
oldest_entry=None,
newest_entry=None,
hit_count=0
)
async def cleanup_expired(self):
"""Remove expired cache entries."""
if not self.enabled or not self._conn:
return
async with self._lock:
try:
await self._conn.execute(
"DELETE FROM search_cache WHERE expires_at < datetime('now')"
)
await self._conn.commit()
except Exception as e:
print(f"[CacheManager] CLEANUP error: {e}", flush=True)
async def close(self):
"""Close the database connection."""
if self._conn:
await self._conn.close()
self._conn = None
# =============================================================================
# Proxy Stats Manager - Track content fetching success/failures per domain
# =============================================================================
class ProxyStatsManager:
"""
Track content fetching statistics per domain.
Records:
- Simple fetch successes/failures
- Proxy fetch successes/failures
- Last attempt timestamps
This data helps identify which domains consistently block proxies.
"""
def __init__(self, db_path: str):
self.db_path = db_path
self._conn = None
async def initialize(self):
"""Initialize the database connection and schema."""
import aiosqlite
import os
os.makedirs(os.path.dirname(self.db_path), exist_ok=True)
self._conn = await aiosqlite.connect(self.db_path)
await self._conn.execute("PRAGMA journal_mode=WAL")
# Create proxy stats table (includes crawl4ai columns)
await self._conn.execute("""
CREATE TABLE IF NOT EXISTS proxy_stats (
domain TEXT PRIMARY KEY,
simple_success INTEGER DEFAULT 0,
simple_failure INTEGER DEFAULT 0,
crawl4ai_success INTEGER DEFAULT 0,
crawl4ai_failure INTEGER DEFAULT 0,
proxy_success INTEGER DEFAULT 0,
proxy_failure INTEGER DEFAULT 0,
skip_domain INTEGER DEFAULT 0,
first_seen TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
last_seen TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)
""")
# Try to add crawl4ai columns if they don't exist (migration for existing DBs)
try:
await self._conn.execute("ALTER TABLE proxy_stats ADD COLUMN crawl4ai_success INTEGER DEFAULT 0")
except Exception:
pass # Column already exists
try:
await self._conn.execute("ALTER TABLE proxy_stats ADD COLUMN crawl4ai_failure INTEGER DEFAULT 0")
except Exception:
pass # Column already exists
# Create index for sorting
await self._conn.execute("""
CREATE INDEX IF NOT EXISTS idx_proxy_stats_failure
ON proxy_stats(proxy_failure DESC)
""")
# Create fetch_errors table for detailed error tracking
await self._conn.execute("""
CREATE TABLE IF NOT EXISTS fetch_errors (
id INTEGER PRIMARY KEY AUTOINCREMENT,
domain TEXT NOT NULL,
fetch_method TEXT NOT NULL,
error_type TEXT NOT NULL,
error_detail TEXT,
url TEXT,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)
""")
# Create index for error queries
await self._conn.execute("""
CREATE INDEX IF NOT EXISTS idx_fetch_errors_domain
ON fetch_errors(domain, error_type)
""")
await self._conn.execute("""
CREATE INDEX IF NOT EXISTS idx_fetch_errors_created
ON fetch_errors(created_at DESC)
""")
await self._conn.commit()
def _extract_domain(self, url: str) -> str:
"""Extract domain from URL."""
from urllib.parse import urlparse
try:
parsed = urlparse(url)
domain = parsed.netloc.lower()
# Remove www. prefix for consistency
if domain.startswith('www.'):
domain = domain[4:]
return domain
except Exception:
return "unknown"
async def record_simple_success(self, url: str):
"""Record a successful simple fetch."""
domain = self._extract_domain(url)
await self._record_stat(domain, "simple_success")
async def record_simple_failure(self, url: str):
"""Record a failed simple fetch."""
domain = self._extract_domain(url)
await self._record_stat(domain, "simple_failure")
async def record_crawl4ai_success(self, url: str):
"""Record a successful crawl4ai fetch."""
domain = self._extract_domain(url)
await self._record_stat(domain, "crawl4ai_success")
async def record_crawl4ai_failure(self, url: str):
"""Record a failed crawl4ai fetch."""
domain = self._extract_domain(url)
await self._record_stat(domain, "crawl4ai_failure")
async def record_proxy_success(self, url: str):
"""Record a successful proxy fetch."""
domain = self._extract_domain(url)
await self._record_stat(domain, "proxy_success")
async def record_proxy_failure(self, url: str):
"""Record a failed proxy fetch."""
domain = self._extract_domain(url)
await self._record_stat(domain, "proxy_failure")
async def record_skip_domain(self, url: str):
"""Record a skipped domain."""
domain = self._extract_domain(url)
await self._record_stat(domain, "skip_domain")
async def _record_stat(self, domain: str, stat_type: str):
"""Record a stat for a domain."""
if not self._conn:
return
try:
# Use upsert to increment counter
await self._conn.execute(f"""
INSERT INTO proxy_stats (domain, {stat_type}, first_seen, last_seen)
VALUES (?, 1, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
ON CONFLICT(domain) DO UPDATE SET
{stat_type} = {stat_type} + 1,
last_seen = CURRENT_TIMESTAMP
""", (domain,))
await self._conn.commit()
except Exception as e:
print(f"[ProxyStats] Error recording stat: {e}", flush=True)
async def get_all_stats(self, limit: int = 100) -> List[dict]:
"""Get all domain stats, sorted by proxy failure rate."""
if not self._conn:
return []
try:
cursor = await self._conn.execute("""
SELECT
domain,
simple_success,
simple_failure,
crawl4ai_success,
crawl4ai_failure,
proxy_success,
proxy_failure,
skip_domain,
first_seen,
last_seen,
CASE
WHEN (proxy_success + proxy_failure) > 0
THEN ROUND(CAST(proxy_failure AS FLOAT) / (proxy_success + proxy_failure) * 100, 1)
ELSE 0
END as proxy_block_rate,
(simple_success + simple_failure + crawl4ai_success + crawl4ai_failure + proxy_success + proxy_failure + skip_domain) as total_attempts
FROM proxy_stats
ORDER BY proxy_failure DESC, total_attempts DESC
LIMIT ?
""", (limit,))
rows = await cursor.fetchall()
return [
{
"domain": row[0],
"simple_success": row[1],
"simple_failure": row[2],
"crawl4ai_success": row[3],
"crawl4ai_failure": row[4],
"proxy_success": row[5],
"proxy_failure": row[6],
"skip_domain": row[7],
"first_seen": row[8],
"last_seen": row[9],
"proxy_block_rate": row[10],
"total_attempts": row[11],
}
for row in rows
]
except Exception as e:
print(f"[ProxyStats] Error getting stats: {e}", flush=True)
return []
async def get_summary(self) -> dict:
"""Get summary statistics."""
if not self._conn:
return {}
try:
cursor = await self._conn.execute("""
SELECT
COUNT(*) as total_domains,
SUM(simple_success) as total_simple_success,
SUM(simple_failure) as total_simple_failure,
SUM(crawl4ai_success) as total_crawl4ai_success,
SUM(crawl4ai_failure) as total_crawl4ai_failure,
SUM(proxy_success) as total_proxy_success,
SUM(proxy_failure) as total_proxy_failure,
SUM(skip_domain) as total_skipped
FROM proxy_stats
""")
row = await cursor.fetchone()
if not row:
return {}
total_simple_attempts = (row[1] or 0) + (row[2] or 0)
total_crawl4ai_attempts = (row[3] or 0) + (row[4] or 0)
total_proxy_attempts = (row[5] or 0) + (row[6] or 0)
return {
"total_domains": row[0] or 0,
"total_simple_success": row[1] or 0,
"total_simple_failure": row[2] or 0,
"total_crawl4ai_success": row[3] or 0,
"total_crawl4ai_failure": row[4] or 0,
"total_proxy_success": row[5] or 0,
"total_proxy_failure": row[6] or 0,
"total_skipped": row[7] or 0,
"simple_success_rate": round((row[1] or 0) / total_simple_attempts * 100, 1) if total_simple_attempts > 0 else 0,
"crawl4ai_success_rate": round((row[3] or 0) / total_crawl4ai_attempts * 100, 1) if total_crawl4ai_attempts > 0 else 0,
"proxy_success_rate": round((row[5] or 0) / total_proxy_attempts * 100, 1) if total_proxy_attempts > 0 else 0,
}
except Exception as e:
print(f"[ProxyStats] Error getting summary: {e}", flush=True)
return {}
async def record_error(self, url: str, fetch_method: str, error_type: str, error_detail: str = None):
"""
Record a fetch error with details.
Args:
url: URL that was being fetched
fetch_method: Which method failed (simple, crawl4ai, proxy)
error_type: Type of error (timeout, forbidden, bot_blocked, connection, extraction, etc.)
error_detail: Additional error details
"""
if not self._conn:
return
domain = self._extract_domain(url)
try:
await self._conn.execute("""
INSERT INTO fetch_errors (domain, fetch_method, error_type, error_detail, url)
VALUES (?, ?, ?, ?, ?)
""", (domain, fetch_method, error_type, error_detail, url))
await self._conn.commit()
except Exception as e:
print(f"[ProxyStats] Error recording error: {e}", flush=True)
async def get_error_summary(self) -> dict:
"""Get aggregated error statistics."""
if not self._conn:
return {}
try:
# Get error counts by type
cursor = await self._conn.execute("""
SELECT
error_type,
fetch_method,
COUNT(*) as count
FROM fetch_errors
GROUP BY error_type, fetch_method
ORDER BY count DESC
""")
rows = await cursor.fetchall()
error_by_type = {}
for row in rows:
error_type = row[0]
method = row[1]
count = row[2]
if error_type not in error_by_type:
error_by_type[error_type] = {"total": 0, "by_method": {}}
error_by_type[error_type]["total"] += count
error_by_type[error_type]["by_method"][method] = count
# Get top domains with errors
cursor = await self._conn.execute("""
SELECT
domain,
COUNT(*) as error_count,
GROUP_CONCAT(DISTINCT error_type) as error_types
FROM fetch_errors
GROUP BY domain
ORDER BY error_count DESC
LIMIT 50
""")
rows = await cursor.fetchall()
top_error_domains = [
{
"domain": row[0],
"error_count": row[1],
"error_types": row[2].split(",") if row[2] else []
}
for row in rows
]
# Get recent errors
cursor = await self._conn.execute("""
SELECT
domain,
fetch_method,
error_type,
error_detail,
url,
datetime(created_at) as created_at
FROM fetch_errors
ORDER BY created_at DESC
LIMIT 100
""")
rows = await cursor.fetchall()
recent_errors = [
{
"domain": row[0],
"fetch_method": row[1],
"error_type": row[2],
"error_detail": row[3],
"url": row[4],
"created_at": row[5]
}
for row in rows
]
# Total error count
cursor = await self._conn.execute("SELECT COUNT(*) FROM fetch_errors")
total = (await cursor.fetchone())[0]
return {
"total_errors": total,
"by_type": error_by_type,
"top_error_domains": top_error_domains,
"recent_errors": recent_errors
}
except Exception as e:
print(f"[ProxyStats] Error getting error summary: {e}", flush=True)
return {}
async def get_domain_errors(self, domain: str, limit: int = 50) -> list:
"""Get errors for a specific domain."""
if not self._conn:
return []
try:
cursor = await self._conn.execute("""
SELECT
fetch_method,
error_type,
error_detail,
url,
datetime(created_at) as created_at
FROM fetch_errors
WHERE domain = ?
ORDER BY created_at DESC
LIMIT ?
""", (domain, limit))
rows = await cursor.fetchall()
return [
{
"fetch_method": row[0],
"error_type": row[1],
"error_detail": row[2],
"url": row[3],
"created_at": row[4]
}
for row in rows
]
except Exception as e:
print(f"[ProxyStats] Error getting domain errors: {e}", flush=True)
return []
async def clear_stats(self):
"""Clear all stats."""
if not self._conn:
return
try:
await self._conn.execute("DELETE FROM proxy_stats")
await self._conn.execute("DELETE FROM fetch_errors")
await self._conn.commit()
except Exception as e:
print(f"[ProxyStats] Error clearing stats: {e}", flush=True)
async def close(self):
"""Close the database connection."""
if self._conn:
await self._conn.close()
self._conn = None
# =============================================================================
# Provider Stats Manager - Track search performance per provider
# =============================================================================
class ProviderStatsManager:
"""
Track search performance statistics per provider.
Records:
- Total searches and success/error counts
- Average search time
- Results returned count
- Content fetch success rate
- Recent search history with detailed metrics
"""
def __init__(self, db_path: str):
self.db_path = db_path
self._conn = None
async def initialize(self):
"""Initialize the database connection and schema."""
import aiosqlite
import os
os.makedirs(os.path.dirname(self.db_path), exist_ok=True)
self._conn = await aiosqlite.connect(self.db_path)
await self._conn.execute("PRAGMA journal_mode=WAL")
# Aggregate stats per provider
await self._conn.execute("""
CREATE TABLE IF NOT EXISTS provider_stats (
provider TEXT PRIMARY KEY,
total_searches INTEGER DEFAULT 0,
successful_searches INTEGER DEFAULT 0,
failed_searches INTEGER DEFAULT 0,
total_results_returned INTEGER DEFAULT 0,
total_results_with_content INTEGER DEFAULT 0,
total_search_time_ms INTEGER DEFAULT 0,
first_seen TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
last_seen TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)
""")
# Individual search logs for detailed analysis
await self._conn.execute("""
CREATE TABLE IF NOT EXISTS search_logs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
provider TEXT NOT NULL,
query TEXT NOT NULL,
status TEXT NOT NULL,
results_count INTEGER DEFAULT 0,
results_with_content INTEGER DEFAULT 0,
search_time_ms REAL DEFAULT 0,
error_message TEXT,
cached INTEGER DEFAULT 0,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)
""")
# Index for faster queries
await self._conn.execute("""
CREATE INDEX IF NOT EXISTS idx_search_logs_provider
ON search_logs(provider, created_at DESC)
""")
await self._conn.commit()
async def record_search(
self,
provider: str,
query: str,
status: str, # 'success', 'error', 'cached'
results_count: int = 0,
results_with_content: int = 0,
search_time_ms: float = 0,
error_message: str = None,
cached: bool = False,
):
"""Record a search attempt with detailed metrics."""
if not self._conn:
return
try:
# Insert individual log
await self._conn.execute("""
INSERT INTO search_logs
(provider, query, status, results_count, results_with_content, search_time_ms, error_message, cached)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
""", (provider, query, status, results_count, results_with_content, search_time_ms, error_message, 1 if cached else 0))
# Update aggregate stats (only for non-cached searches)
if not cached:
is_success = status == 'success'
await self._conn.execute("""
INSERT INTO provider_stats
(provider, total_searches, successful_searches, failed_searches,
total_results_returned, total_results_with_content, total_search_time_ms)
VALUES (?, 1, ?, ?, ?, ?, ?)
ON CONFLICT(provider) DO UPDATE SET
total_searches = total_searches + 1,
successful_searches = successful_searches + ?,
failed_searches = failed_searches + ?,
total_results_returned = total_results_returned + ?,
total_results_with_content = total_results_with_content + ?,
total_search_time_ms = total_search_time_ms + ?,
last_seen = CURRENT_TIMESTAMP
""", (
provider,
1 if is_success else 0,
0 if is_success else 1,
results_count,
results_with_content,
int(search_time_ms),
# For update
1 if is_success else 0,
0 if is_success else 1,
results_count,
results_with_content,
int(search_time_ms),
))
await self._conn.commit()
except Exception as e:
print(f"[ProviderStats] Error recording search: {e}", flush=True)
async def get_provider_stats(self) -> List[dict]:
"""Get aggregate statistics per provider."""
if not self._conn:
return []
try:
cursor = await self._conn.execute("""
SELECT
provider,
total_searches,
successful_searches,
failed_searches,
total_results_returned,
total_results_with_content,
total_search_time_ms,
first_seen,
last_seen,
CASE WHEN total_searches > 0
THEN ROUND(CAST(successful_searches AS FLOAT) / total_searches * 100, 1)
ELSE 0 END as success_rate,
CASE WHEN total_searches > 0
THEN ROUND(CAST(total_search_time_ms AS FLOAT) / total_searches, 0)
ELSE 0 END as avg_search_time_ms,
CASE WHEN total_results_returned > 0
THEN ROUND(CAST(total_results_with_content AS FLOAT) / total_results_returned * 100, 1)
ELSE 0 END as content_success_rate,
CASE WHEN total_searches > 0
THEN ROUND(CAST(total_results_returned AS FLOAT) / total_searches, 1)
ELSE 0 END as avg_results_per_search
FROM provider_stats
ORDER BY total_searches DESC
""")
rows = await cursor.fetchall()
return [
{
"provider": row[0],
"total_searches": row[1],
"successful_searches": row[2],
"failed_searches": row[3],
"total_results_returned": row[4],
"total_results_with_content": row[5],
"total_search_time_ms": row[6],
"first_seen": row[7],
"last_seen": row[8],
"success_rate": row[9],
"avg_search_time_ms": row[10],
"content_success_rate": row[11],
"avg_results_per_search": row[12],
}
for row in rows
]
except Exception as e:
print(f"[ProviderStats] Error getting stats: {e}", flush=True)
return []
async def get_recent_searches(self, provider: str = None, limit: int = 50) -> List[dict]:
"""Get recent search logs."""
if not self._conn:
return []
try:
if provider:
cursor = await self._conn.execute("""
SELECT
id, provider, query, status, results_count, results_with_content,
search_time_ms, error_message, cached, datetime(created_at) as created_at
FROM search_logs
WHERE provider = ?
ORDER BY created_at DESC
LIMIT ?
""", (provider, limit))
else:
cursor = await self._conn.execute("""
SELECT
id, provider, query, status, results_count, results_with_content,
search_time_ms, error_message, cached, datetime(created_at) as created_at
FROM search_logs
ORDER BY created_at DESC
LIMIT ?
""", (limit,))
rows = await cursor.fetchall()
return [
{
"id": row[0],
"provider": row[1],
"query": row[2],
"status": row[3],
"results_count": row[4],
"results_with_content": row[5],
"search_time_ms": row[6],
"error_message": row[7],
"cached": bool(row[8]),
"created_at": row[9],
}
for row in rows
]
except Exception as e:
print(f"[ProviderStats] Error getting recent searches: {e}", flush=True)
return []
async def get_summary(self) -> dict:
"""Get overall summary statistics."""
if not self._conn:
return {}
try:
cursor = await self._conn.execute("""
SELECT
COUNT(DISTINCT provider) as total_providers,
SUM(total_searches) as total_searches,
SUM(successful_searches) as total_successful,
SUM(failed_searches) as total_failed,
SUM(total_results_returned) as total_results,
SUM(total_results_with_content) as total_with_content,
SUM(total_search_time_ms) as total_time_ms
FROM provider_stats
""")
row = await cursor.fetchone()
if not row or not row[1]:
return {
"total_providers": 0,
"total_searches": 0,
"total_successful": 0,
"total_failed": 0,
"success_rate": 0,
"total_results": 0,
"total_with_content": 0,
"content_rate": 0,
"avg_search_time_ms": 0,
}
total_searches = row[1] or 0
total_results = row[4] or 0
return {
"total_providers": row[0] or 0,
"total_searches": total_searches,
"total_successful": row[2] or 0,
"total_failed": row[3] or 0,
"success_rate": round((row[2] or 0) / total_searches * 100, 1) if total_searches > 0 else 0,
"total_results": total_results,
"total_with_content": row[5] or 0,
"content_rate": round((row[5] or 0) / total_results * 100, 1) if total_results > 0 else 0,
"avg_search_time_ms": round((row[6] or 0) / total_searches, 0) if total_searches > 0 else 0,
}
except Exception as e:
print(f"[ProviderStats] Error getting summary: {e}", flush=True)
return {}
async def clear_stats(self):
"""Clear all stats."""
if not self._conn:
return
try:
await self._conn.execute("DELETE FROM provider_stats")
await self._conn.execute("DELETE FROM search_logs")
await self._conn.commit()
except Exception as e:
print(f"[ProviderStats] Error clearing stats: {e}", flush=True)
async def close(self):
"""Close the database connection."""
if self._conn:
await self._conn.close()
self._conn = None
# =============================================================================
# Global State
# =============================================================================
provider_registry: Optional[ProviderRegistry] = None
cache_manager: Optional[CacheManager] = None
proxy_stats_manager: Optional[ProxyStatsManager] = None
provider_stats_manager: Optional[ProviderStatsManager] = None
# =============================================================================
# Application Lifecycle
# =============================================================================
@asynccontextmanager
async def lifespan(app: FastAPI):
"""Application lifecycle manager."""
global provider_registry, cache_manager, proxy_stats_manager, provider_stats_manager
print("🔍 Web Search Service starting...")
print(f" Default provider: {config.default_provider}")
print(f" Cache enabled: {config.cache_enabled}")
print(f" Cache TTL: {config.cache_ttl_hours} hours")
print(f" Max concurrent: {config.max_concurrent}")
print(f" Max queue size: {config.max_queue_size} (backpressure limit)")
# Initialize cache (with enabled flag from config)
cache_manager = CacheManager(config.database_path, config.cache_ttl_hours, enabled=config.cache_enabled)
await cache_manager.initialize()
if config.cache_enabled:
print(f" Cache initialized: {config.database_path}")
else:
print(f" Cache DISABLED (set WEBSEARCH_CACHE_ENABLED=true to enable)")
# Initialize proxy stats tracking
proxy_stats_manager = ProxyStatsManager(config.database_path)
await proxy_stats_manager.initialize()
print(f" Proxy stats tracking initialized")
# Initialize provider stats tracking
provider_stats_manager = ProviderStatsManager(config.database_path)
await provider_stats_manager.initialize()
print(f" Provider stats tracking initialized")
# Set up content fetcher stats callbacks (includes crawl4ai and error recording)
from providers.content_fetcher import set_stats_callbacks
set_stats_callbacks(
simple_success=proxy_stats_manager.record_simple_success,
simple_failure=proxy_stats_manager.record_simple_failure,
crawl4ai_success=proxy_stats_manager.record_crawl4ai_success,
crawl4ai_failure=proxy_stats_manager.record_crawl4ai_failure,
proxy_success=proxy_stats_manager.record_proxy_success,
proxy_failure=proxy_stats_manager.record_proxy_failure,
skip_domain=proxy_stats_manager.record_skip_domain,
record_error=proxy_stats_manager.record_error,
)
print(f" Content fetcher stats callbacks registered (with crawl4ai and error tracking)")
# Initialize providers
provider_registry = ProviderRegistry(config)
await provider_registry.initialize()
available = [p.name for p in provider_registry.get_all_providers_info() if p.available]
print(f" Available providers: {', '.join(available) or 'none'}")
print(" ✅ Web Search Service ready")
yield
# Cleanup
print("👋 Web Search Service shutting down...")
if cache_manager:
await cache_manager.close()
if proxy_stats_manager:
await proxy_stats_manager.close()
if provider_stats_manager:
await provider_stats_manager.close()
if provider_registry:
await provider_registry.close_all()
# =============================================================================
# FastAPI Application
# =============================================================================
app = FastAPI(
title="Web Search Service",
description="Centralized web search proxy with provider abstraction and caching",
version="1.0.0",
lifespan=lifespan
)
# CORS middleware
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# Request ID middleware for tracing
app.add_middleware(RequestIDMiddleware)
# =============================================================================
# Endpoints
# =============================================================================
@app.get("/health", response_model=HealthResponse)
async def health_check():
"""Health check endpoint with provider status."""
providers = []
if provider_registry:
providers = provider_registry.get_all_providers_info()
return HealthResponse(
status="healthy",
service="websearch",
version="1.0.0",
providers=providers,
cache_enabled=cache_manager is not None,
default_provider=config.default_provider
)
@app.get("/providers", response_model=List[ProviderInfo])
async def list_providers():
"""List all available search providers."""
if not provider_registry:
raise HTTPException(status_code=503, detail="Service not initialized")
return provider_registry.get_all_providers_info()
@app.get("/providers/google/quota")
async def get_google_quota():
"""
Get Google Custom Search API quota status.
Returns remaining queries for today and the daily limit.
"""
if not provider_registry:
raise HTTPException(status_code=503, detail="Service not initialized")
try:
provider, _ = provider_registry.get_provider("google")
remaining = await provider.get_remaining_quota()
return {
"provider": "google",
"daily_limit": provider.daily_limit,
"used_today": provider.daily_limit - remaining,
"remaining_today": remaining,
"fetch_content_enabled": provider.fetch_content,
}
except Exception as e:
raise HTTPException(status_code=404, detail=f"Google provider not available: {e}")
@app.get("/cache/stats", response_model=CacheStatsResponse)
async def get_cache_stats():
"""Get cache statistics."""
if not cache_manager:
raise HTTPException(status_code=503, detail="Cache not initialized")
return await cache_manager.get_stats()
@app.post("/cache/cleanup")
async def cleanup_cache():
"""Manually trigger cache cleanup of expired entries."""
if not cache_manager:
raise HTTPException(status_code=503, detail="Cache not initialized")
await cache_manager.cleanup_expired()
return {"status": "ok", "message": "Expired entries removed"}
@app.get("/cache/entries")
async def get_cache_entries(
provider: Optional[str] = Query(None, description="Filter by provider"),
search: Optional[str] = Query(None, description="Search query text (case-insensitive)"),
limit: int = Query(50, ge=1, le=200, description="Max entries to return"),
offset: int = Query(0, ge=0, description="Offset for pagination"),
):
"""Get cached search entries with optional filtering and search."""
if not cache_manager or not cache_manager._conn:
raise HTTPException(status_code=503, detail="Cache not initialized")
# Build query with WHERE conditions
query = """
SELECT
cache_key,
query,
provider,
start_date,
end_date,
max_results,
hit_count,
datetime(created_at) as created_at,
datetime(expires_at) as expires_at
FROM search_cache
"""
params = []
conditions = []
if provider:
conditions.append("provider = ?")
params.append(provider)
if search:
# Case-insensitive search in query text
conditions.append("LOWER(query) LIKE ?")
params.append(f"%{search.lower()}%")
if conditions:
query += " WHERE " + " AND ".join(conditions)
query += " ORDER BY created_at DESC LIMIT ? OFFSET ?"
params.extend([limit, offset])
cursor = await cache_manager._conn.execute(query, params)
rows = await cursor.fetchall()
entries = []
for row in rows:
entries.append({
"cache_key": row[0],
"query": row[1],
"provider": row[2],
"start_date": row[3],
"end_date": row[4],
"max_results": row[5],
"hit_count": row[6],
"created_at": row[7],
"expires_at": row[8],
})
# Get total count with same filters
count_query = "SELECT COUNT(*) FROM search_cache"
count_params = []
if conditions:
count_query += " WHERE " + " AND ".join(conditions)
# Re-add filter params (without limit/offset)
if provider:
count_params.append(provider)
if search:
count_params.append(f"%{search.lower()}%")
cursor = await cache_manager._conn.execute(count_query, count_params)
total = (await cursor.fetchone())[0]
return {
"entries": entries,
"total": total,
"limit": limit,
"offset": offset,
"provider_filter": provider,
"search_filter": search,
}
@app.get("/cache/providers")
async def get_cache_providers():
"""Get cache statistics grouped by provider."""
if not cache_manager or not cache_manager._conn:
raise HTTPException(status_code=503, detail="Cache not initialized")
cursor = await cache_manager._conn.execute("""
SELECT
provider,
COUNT(*) as entry_count,
COUNT(DISTINCT query) as unique_queries,
SUM(hit_count) as total_hits,
MIN(created_at) as oldest_entry,
MAX(created_at) as newest_entry
FROM search_cache
GROUP BY provider
ORDER BY entry_count DESC
""")
rows = await cursor.fetchall()
providers = []
for row in rows:
providers.append({
"provider": row[0],
"entry_count": row[1],
"unique_queries": row[2],
"total_hits": row[3] or 0,
"oldest_entry": row[4],
"newest_entry": row[5],
})
return {"providers": providers}
@app.get("/cache/entry/{cache_key}")
async def get_cache_entry_detail(cache_key: str):
"""
Get full cached search results by cache key.
Returns the complete results including snippets and raw_content.
"""
if not cache_manager or not cache_manager._conn:
raise HTTPException(status_code=503, detail="Cache not initialized")
cursor = await cache_manager._conn.execute(
"""
SELECT
cache_key,
query,
provider,
start_date,
end_date,
max_results,
results,
hit_count,
datetime(created_at) as created_at,
datetime(expires_at) as expires_at
FROM search_cache
WHERE cache_key = ?
""",
(cache_key,)
)
row = await cursor.fetchone()
if not row:
raise HTTPException(status_code=404, detail="Cache entry not found")
# Parse results from JSON
results_data = json.loads(row[6])
return {
"cache_key": row[0],
"query": row[1],
"provider": row[2],
"start_date": row[3],
"end_date": row[4],
"max_results": row[5],
"results": results_data,
"result_count": len(results_data),
"hit_count": row[7],
"created_at": row[8],
"expires_at": row[9],
}
# =============================================================================
# Proxy Stats Endpoints - Track content fetching success/failures per domain
# =============================================================================
@app.get("/proxy-stats")
async def get_proxy_stats(limit: int = 100):
"""
Get proxy statistics per domain.
Returns domains sorted by proxy failure count, showing:
- Simple fetch success/failure counts
- Proxy fetch success/failure counts
- Proxy block rate (% of proxy attempts that failed)
This helps identify domains that consistently block residential proxies.
"""
if not proxy_stats_manager:
raise HTTPException(status_code=503, detail="Proxy stats not initialized")
stats = await proxy_stats_manager.get_all_stats(limit)
summary = await proxy_stats_manager.get_summary()
return {
"summary": summary,
"domains": stats,
}
@app.get("/proxy-stats/summary")
async def get_proxy_stats_summary():
"""Get summary of proxy statistics."""
if not proxy_stats_manager:
raise HTTPException(status_code=503, detail="Proxy stats not initialized")
return await proxy_stats_manager.get_summary()
@app.post("/proxy-stats/clear")
async def clear_proxy_stats():
"""Clear all proxy statistics."""
if not proxy_stats_manager:
raise HTTPException(status_code=503, detail="Proxy stats not initialized")
await proxy_stats_manager.clear_stats()
return {"status": "cleared"}
# =============================================================================
# Fetch Error Endpoints - Detailed error tracking
# =============================================================================
@app.get("/fetch-errors")
async def get_fetch_errors():
"""
Get detailed fetch error statistics.
Returns:
- Error counts by type (timeout, forbidden, bot_blocked, etc.)
- Top domains with errors
- Recent error log
This helps identify:
- Which error types are most common
- Which domains are most problematic
- Recent failures for debugging
"""
if not proxy_stats_manager:
raise HTTPException(status_code=503, detail="Proxy stats not initialized")
return await proxy_stats_manager.get_error_summary()
@app.get("/fetch-errors/domain/{domain}")
async def get_domain_fetch_errors(domain: str, limit: int = Query(50, ge=1, le=200)):
"""
Get fetch errors for a specific domain.
Args:
domain: Domain to get errors for (without www.)
limit: Maximum number of errors to return
Returns:
List of recent errors for the domain
"""
if not proxy_stats_manager:
raise HTTPException(status_code=503, detail="Proxy stats not initialized")
return {
"domain": domain,
"errors": await proxy_stats_manager.get_domain_errors(domain, limit)
}
# =============================================================================
# Provider Stats Endpoints - Track search performance per provider
# =============================================================================
@app.get("/provider-stats")
async def get_provider_stats():
"""
Get aggregate statistics per search provider.
Returns:
- Total searches, success/error counts
- Average search time
- Results returned and content success rates
"""
if not provider_stats_manager:
raise HTTPException(status_code=503, detail="Provider stats not initialized")
providers = await provider_stats_manager.get_provider_stats()
summary = await provider_stats_manager.get_summary()
return {
"summary": summary,
"providers": providers,
}
@app.get("/provider-stats/summary")
async def get_provider_stats_summary():
"""Get summary of provider statistics."""
if not provider_stats_manager:
raise HTTPException(status_code=503, detail="Provider stats not initialized")
return await provider_stats_manager.get_summary()
@app.get("/provider-stats/recent")
async def get_recent_searches(
provider: Optional[str] = Query(None, description="Filter by provider"),
limit: int = Query(50, ge=1, le=200, description="Max entries"),
):
"""Get recent search logs with detailed metrics."""
if not provider_stats_manager:
raise HTTPException(status_code=503, detail="Provider stats not initialized")
return await provider_stats_manager.get_recent_searches(provider, limit)
@app.post("/provider-stats/clear")
async def clear_provider_stats():
"""Clear all provider statistics."""
if not provider_stats_manager:
raise HTTPException(status_code=503, detail="Provider stats not initialized")
await provider_stats_manager.clear_stats()
return {"status": "cleared"}
# =============================================================================
# Queue Status & Blocked Domains Endpoints
# =============================================================================
@app.get("/blocked-domains")
async def get_blocked_domains():
"""
Get list of domains that are blocked from proxy fallback.
These domains consistently block the Brightdata proxy (100% block rate).
Simple fetch and crawl4ai are still attempted, but proxy fallback is skipped.
"""
from providers.content_fetcher import AsyncContentFetcher
return {
"skip_domains": AsyncContentFetcher.SKIP_DOMAINS,
"proxy_blocked_domains": AsyncContentFetcher.PROXY_BLOCKED_DOMAINS,
"note": "skip_domains are not fetched at all; proxy_blocked_domains skip proxy fallback only"
}
@app.get("/queue-status")
async def get_queue_status():
"""
Get current queue status for all providers.
Returns queue waiting count and capacity for each provider.
Useful for monitoring and implementing client-side backoff.
"""
if not provider_registry:
raise HTTPException(status_code=503, detail="Service not initialized")
status = {}
for name, provider in provider_registry._providers.items():
info = provider.get_info()
status[name] = {
"queue_waiting": info.get("queue_waiting", 0),
"queue_capacity": info.get("queue_capacity", 0),
"available": info.get("available", False),
}
return {
"providers": status,
"total_waiting": sum(s["queue_waiting"] for s in status.values()),
"total_capacity": sum(s["queue_capacity"] for s in status.values()),
}
# =============================================================================
# Search Endpoint
# =============================================================================
@app.post("/search", response_model=BatchSearchResponse)
async def search(request: SearchRequest):
"""
Execute web search queries.
Supports multiple queries in a single request, with optional provider selection,
date filtering, and caching.
Filter modes:
- 'none': No post-processing (default)
- 'heuristic': Fast URL-based filtering (removes results with years in URL outside range)
- 'llm': AI-powered content analysis (slower but more accurate)
"""
import time
if not provider_registry:
raise HTTPException(status_code=503, detail="Service not initialized")
# Validate filter mode
valid_modes = ["none", "heuristic", "llm"]
if request.filter_mode not in valid_modes:
raise HTTPException(
status_code=400,
detail=f"Invalid filter_mode. Must be one of: {valid_modes}"
)
start_time = time.time()
# Get provider
provider, provider_name = provider_registry.get_provider(request.provider)
# Log request details
date_info = ""
if request.start_date or request.end_date:
date_info = f", dates: {request.start_date or 'any'}{request.end_date or 'any'}"
log(f"Search: provider={provider_name}, queries={len(request.queries)}, max={request.max_results}{date_info}")
# Initialize date filter if needed
date_filter = None
if request.filter_mode != "none" and request.start_date and request.end_date:
date_filter = DateFilter(request.start_date, request.end_date)
searches = []
cache_hits = 0
cache_misses = 0
# Aggregate filter stats
total_before = 0
total_after = 0
all_filtered_urls = []
for query in request.queries:
query_start = time.time()
cached = False
results = []
# Check cache first
if request.use_cache and cache_manager:
cached_results = await cache_manager.get(
query=query,
provider=provider_name,
start_date=request.start_date,
end_date=request.end_date,
max_results=request.max_results
)
if cached_results is not None:
results = cached_results
cached = True
cache_hits += 1
# If not cached, search
if not cached:
cache_misses += 1
try:
results = await provider.search(
query=query,
max_results=request.max_results,
start_date=request.start_date,
end_date=request.end_date
)
# Cache results (before filtering - we cache raw results)
if request.use_cache and cache_manager:
await cache_manager.set(
query=query,
provider=provider_name,
start_date=request.start_date,
end_date=request.end_date,
max_results=request.max_results,
results=results
)
except QueueFullError as qfe:
# Backpressure: queue is full, reject immediately
raise HTTPException(
status_code=503,
detail=f"Service overloaded: {str(qfe)}. Reduce request rate or try again later.",
headers={"Retry-After": "10"} # Suggest retry after 10 seconds
)
except Exception as e:
error_str = str(e).lower()
# Determine if this is a transient/retryable error
is_transient = any(indicator in error_str for indicator in [
'timeout', 'econnrefused', 'connection', 'socket hang up',
'proxy request failed', '504', '502', '503'
])
# Retry for Brightdata transient errors with exponential backoff
if provider_name == "brightdata":
retry_delay = 3 if is_transient else 2 # Longer delay for timeouts
print(f"[Brightdata] First attempt failed: {e}, retrying after {retry_delay}s delay...", flush=True)
await asyncio.sleep(retry_delay)
try:
results = await provider.search(
query=query,
max_results=request.max_results,
start_date=request.start_date,
end_date=request.end_date
)
if request.use_cache and cache_manager:
await cache_manager.set(
query=query,
provider=provider_name,
start_date=request.start_date,
end_date=request.end_date,
max_results=request.max_results,
results=results
)
except Exception as retry_error:
# Record error stats
if provider_stats_manager:
await provider_stats_manager.record_search(
provider=provider_name,
query=query,
status='error',
search_time_ms=(time.time() - query_start) * 1000,
error_message=str(retry_error),
)
# Return 503 for transient errors (client should retry)
# Return 500 for permanent failures
status_code = 503 if is_transient else 500
raise HTTPException(
status_code=status_code,
detail=f"Search failed for query '{query}' after retry: {str(retry_error)}",
headers={"Retry-After": "15"} if status_code == 503 else None
)
else:
# Record error stats for non-Brightdata providers
if provider_stats_manager:
await provider_stats_manager.record_search(
provider=provider_name,
query=query,
status='error',
search_time_ms=(time.time() - query_start) * 1000,
error_message=str(e),
)
status_code = 503 if is_transient else 500
raise HTTPException(
status_code=status_code,
detail=f"Search failed for query '{query}': {str(e)}",
headers={"Retry-After": "10"} if status_code == 503 else None
)
# Apply date filter if enabled
filtered_info = []
results_before_filter = len(results)
total_before += results_before_filter
if date_filter and results:
print(f"📋 [FILTER] Applying {request.filter_mode} filter for query: '{query}'")
print(f" Date range: {request.start_date}{request.end_date}")
print(f" Results before filtering: {results_before_filter}")
if request.filter_mode == "heuristic":
results, filtered_info = date_filter.filter_heuristic(results)
elif request.filter_mode == "llm":
# Get vLLM URL from environment
import os
vllm_url = os.getenv("VLLM_URL") or os.getenv("OPENAI_API_BASE")
results, filtered_info = await date_filter.filter_llm(results, vllm_url)
# Log detailed filtering results
results_after_filter = len(results)
filtered_count = results_before_filter - results_after_filter
print(f" Results after filtering: {results_after_filter}")
print(f" Filtered out: {filtered_count}")
if filtered_info:
print(f" 🚫 Filtered URLs and reasons:")
for f in filtered_info:
print(f" - {f['url']}")
print(f" Reason: {f['reason']}")
if results:
print(f" ✅ Kept URLs:")
for r in results[:5]: # Show first 5
print(f" - {r.url}")
if len(results) > 5:
print(f" ... and {len(results) - 5} more")
all_filtered_urls.extend([f["url"] for f in filtered_info])
total_after += len(results)
query_time = (time.time() - query_start) * 1000
# Record provider stats
# Setting this to False to avoid recording stats for now
_stats_enabled = False
if _stats_enabled and provider_stats_manager:
results_with_content = sum(1 for r in results if r.raw_content)
await provider_stats_manager.record_search(
provider=provider_name,
query=query,
status='success',
results_count=len(results),
results_with_content=results_with_content,
search_time_ms=query_time,
cached=cached,
)
# Calculate content fetch time from individual results
content_fetch_time_ms = 0.0
result_responses = []
for r in results:
result_dict = r.to_dict()
result_responses.append(SearchResultResponse(**result_dict))
# Sum up fetch times (use max as they run in parallel, but track actual)
if r.fetch_time_ms:
content_fetch_time_ms = max(content_fetch_time_ms, r.fetch_time_ms)
searches.append(SearchResponse(
query=query,
provider=provider_name,
results=result_responses,
cached=cached,
search_time_ms=query_time,
content_fetch_time_ms=content_fetch_time_ms,
total_time_ms=query_time + content_fetch_time_ms
))
total_time = (time.time() - start_time) * 1000
# Build filter stats
filter_stats = None
if request.filter_mode != "none":
filter_stats = FilterStats(
mode=request.filter_mode,
total_before=total_before,
total_after=total_after,
filtered_count=total_before - total_after,
filtered_urls=all_filtered_urls
)
return BatchSearchResponse(
searches=searches,
total_time_ms=total_time,
cache_hits=cache_hits,
cache_misses=cache_misses,
filter_stats=filter_stats
)
# =============================================================================
# Main
# =============================================================================
if __name__ == "__main__":
import uvicorn
uvicorn.run(
"api:app",
host=config.host,
port=config.port,
reload=False,
log_level="info"
)