release: bump version to 0.51.0

This commit is contained in:
rayd1o
2026-05-11 09:49:08 +08:00
parent 455b8360d0
commit 1cb51b1172
52 changed files with 4440 additions and 279 deletions

View File

@@ -20,7 +20,11 @@ from app.services.bgp_collector_locations import (
)
from app.services.bgp_collectors import build_bgp_collector_coverage
from app.services.ai_client import get_ai_provider_client
from app.services.location.llm_fallback import collect_llm_location_fallback_candidate
from app.api.v1.settings import get_web_search_client
from app.services.location.llm_fallback import (
collect_llm_location_fallback_candidate,
collect_location_search_evidence,
)
router = APIRouter()
@@ -320,21 +324,33 @@ async def collect_bgp_collector_location(
country=country,
operator=operator,
)
llm_result = None
try:
web_search_client = await get_web_search_client(db)
search_result = await collect_location_search_evidence(
web_search_client=web_search_client,
query=query,
entity_type="bgp_collector",
)
attempted_queries = [*attempted_queries, *search_result.attempted_queries]
if not search_result.evidence:
llm_failure_reason = search_result.failure_reason
raise RuntimeError(search_result.failure_reason or "no WebSearch evidence")
provider_client = await get_ai_provider_client(db)
llm_result = await collect_llm_location_fallback_candidate(
provider_client=provider_client,
query=query,
entity_type="bgp_collector",
attempted_queries=attempted_queries,
search_evidence=search_result.evidence,
)
except Exception as exc:
llm_result = None
llm_failure_reason = f"LLM location factcheck unavailable: {exc}"
attempted_queries = [
*attempted_queries,
f"llm_factcheck:bgp_collector:{collector_id or 'unknown'}",
]
if llm_failure_reason is None:
llm_failure_reason = f"LLM location factcheck unavailable: {exc}"
attempted_queries = [
*attempted_queries,
f"llm_factcheck:bgp_collector:{collector_id or 'unknown'}",
]
if llm_result is not None:
attempted_queries = [*attempted_queries, *llm_result.attempted_queries]
candidates = llm_result.candidates

View File

@@ -1,5 +1,6 @@
from copy import deepcopy
from datetime import UTC, datetime
import os
from pathlib import Path
from typing import Optional
@@ -38,6 +39,16 @@ from app.services.datasource_connectivity import (
save_connectivity_success,
)
from app.services.ai_client import AIProviderClient, get_ai_provider_client
from app.services.ai_tools.schemas import WebSearchConfig, WebSearchProviderConfig
from app.services.ai_tools.web_search import (
WebSearchClient,
WebSearchConfigurationError,
WebSearchError,
get_web_search_provider_preset,
list_web_search_provider_presets,
normalize_web_search_provider,
provider_defaults as web_search_provider_defaults,
)
from app.services.llm_provider_catalog import (
FALLBACK_LLM_PROVIDER_PRESETS,
get_fallback_llm_provider_preset,
@@ -78,7 +89,12 @@ DEFAULT_SETTINGS = {
"providers": {},
"timeout_seconds": 60,
"retry_attempts": 2,
}
},
"web_search": {
"enabled": False,
"default_provider": "tavily",
"providers": {},
},
},
}
@@ -161,9 +177,31 @@ class BarentsWatchIntegrationUpdate(BaseModel):
clear_client_secret: bool = False
class WebSearchIntegrationUpdate(BaseModel):
enabled: bool = False
default_provider: Optional[str] = None
provider: str = Field(default="tavily", max_length=80)
base_url: str = Field(default="", max_length=500)
api_key: Optional[str] = None
max_results: int = Field(default=5, ge=1, le=20)
timeout_seconds: int = Field(default=20, ge=3, le=120)
endpoint_path: str = Field(default="", max_length=200)
search_depth: str = Field(default="basic", max_length=40)
engine: str = Field(default="google", max_length=80)
include_answer: bool = False
include_raw_content: bool = False
include_text: bool = False
categories: str = Field(default="general", max_length=120)
engines: list[str] = Field(default_factory=list)
search_path: str = Field(default="", max_length=200)
scrape_path: str = Field(default="", max_length=200)
scrape_formats: list[str] = Field(default_factory=lambda: ["markdown"])
class ExternalIntegrationsUpdate(BaseModel):
ai_provider: AIProviderIntegrationUpdate
barentswatch: BarentsWatchIntegrationUpdate
web_search: WebSearchIntegrationUpdate | None = None
def merge_with_defaults(category: str, payload: Optional[dict]) -> dict:
@@ -217,6 +255,11 @@ async def save_setting_payload(db: AsyncSession, category: str, payload: dict) -
AI_PROVIDER_ENV_FILE = Path(__file__).resolve().parents[4] / "aiprovider" / ".env"
WEB_SEARCH_ENV_FILES = (
Path(__file__).resolve().parents[4] / ".env",
Path(__file__).resolve().parents[3] / ".env",
AI_PROVIDER_ENV_FILE,
)
def _mask_secret(value: Optional[str], source: str = "") -> dict:
@@ -270,6 +313,33 @@ def _resolve_env_secret(*names: str) -> tuple[str, str]:
return "", ""
def _read_web_search_env_files() -> dict[str, str]:
values: dict[str, str] = {}
for path in WEB_SEARCH_ENV_FILES:
if not path.exists():
continue
values.update({
key: str(value)
for key, value in dotenv_values(path).items()
if value is not None
})
return values
def _resolve_web_search_env_secret(*names: str) -> tuple[str, str]:
env_file_values = _read_web_search_env_files()
for name in names:
if not name:
continue
value = env_file_values.get(name)
if value:
return value, "env_file"
value = os.environ.get(name)
if value:
return value, "env"
return "", ""
def _provider_defaults(provider: str) -> dict:
preset = _get_provider_preset(provider)
return {
@@ -438,6 +508,160 @@ def _runtime_config_from_ai_payload(ai_payload: dict) -> dict:
}
def _web_search_provider_defaults(provider: str) -> dict:
return web_search_provider_defaults(provider).model_dump()
def _normalize_web_search_payload(web_search_payload: dict | None) -> dict:
raw = dict(web_search_payload or {})
default_provider = normalize_web_search_provider(
raw.get("default_provider") or raw.get("provider")
)
providers = {
normalize_web_search_provider(provider): dict(config or {})
for provider, config in (raw.get("providers") or {}).items()
if provider
}
legacy_fields = {
key: raw.get(key)
for key in (
"base_url",
"api_key",
"max_results",
"timeout_seconds",
"endpoint_path",
"search_depth",
"engine",
"include_answer",
"include_raw_content",
"include_text",
"categories",
"engines",
"search_path",
"scrape_path",
"scrape_formats",
)
if raw.get(key) not in (None, "")
}
if legacy_fields:
providers[default_provider] = {
**providers.get(default_provider, {}),
**legacy_fields,
}
normalized_providers: dict[str, dict] = {}
for provider, config in providers.items():
provider_id = normalize_web_search_provider(provider)
normalized_providers[provider_id] = {
**_web_search_provider_defaults(provider_id),
**dict(config or {}),
"provider": provider_id,
}
if default_provider not in normalized_providers:
normalized_providers[default_provider] = _web_search_provider_defaults(default_provider)
return {
"enabled": bool(raw.get("enabled", False)),
"default_provider": default_provider,
"providers": normalized_providers,
}
def _resolve_web_search_api_key(provider: str, provider_config: dict) -> tuple[str, str]:
saved_key = provider_config.get("api_key") or ""
if saved_key:
return str(saved_key), "runtime"
preset = get_web_search_provider_preset(provider)
return _resolve_web_search_env_secret(preset.get("api_key_env") or "", "WEB_SEARCH_API_KEY")
def _build_web_search_payload(
current_payload: dict,
update: WebSearchIntegrationUpdate | None,
) -> dict:
current_web_search = _normalize_web_search_payload(current_payload.get("web_search") or {})
if update is None:
return current_web_search
provider_id = normalize_web_search_provider(update.default_provider or update.provider)
current_providers = {
provider: dict(config or {})
for provider, config in current_web_search.get("providers", {}).items()
}
current_provider = current_providers.get(provider_id) or _web_search_provider_defaults(provider_id)
current_key, current_key_source = _resolve_web_search_api_key(provider_id, current_provider)
current_key_preview = _mask_secret(current_key, current_key_source)["preview"]
provider_payload = {
**_web_search_provider_defaults(provider_id),
**current_provider,
"provider": provider_id,
"base_url": update.base_url.strip()
or current_provider.get("base_url")
or _web_search_provider_defaults(provider_id).get("base_url")
or "",
"max_results": update.max_results,
"timeout_seconds": update.timeout_seconds,
"endpoint_path": update.endpoint_path.strip() or current_provider.get("endpoint_path") or "",
"search_depth": update.search_depth.strip() or "basic",
"engine": update.engine.strip() or "google",
"include_answer": update.include_answer,
"include_raw_content": update.include_raw_content,
"include_text": update.include_text,
"categories": update.categories.strip() or "general",
"engines": update.engines,
"search_path": update.search_path.strip() or current_provider.get("search_path") or "",
"scrape_path": update.scrape_path.strip() or current_provider.get("scrape_path") or "",
"scrape_formats": update.scrape_formats or ["markdown"],
}
if not _is_secret_placeholder(update.api_key, current_key_preview):
provider_payload["api_key"] = str(update.api_key).strip()
elif current_provider.get("api_key"):
provider_payload["api_key"] = current_provider.get("api_key") or ""
else:
provider_payload["api_key"] = ""
current_providers[provider_id] = provider_payload
return {
"enabled": update.enabled,
"default_provider": provider_id,
"providers": current_providers,
}
def _runtime_config_from_web_search_payload(web_search_payload: dict) -> WebSearchConfig:
normalized = _normalize_web_search_payload(web_search_payload)
provider_id = normalized["default_provider"]
provider_config = normalized["providers"].get(provider_id) or _web_search_provider_defaults(provider_id)
api_key, _source = _resolve_web_search_api_key(provider_id, provider_config)
provider_models = {
provider: WebSearchProviderConfig(**{
**config,
"api_key": (
api_key if provider == provider_id else _resolve_web_search_api_key(provider, config)[0]
),
})
for provider, config in normalized["providers"].items()
}
return WebSearchConfig(
enabled=normalized["enabled"],
default_provider=provider_id,
provider=provider_id,
providers=provider_models,
)
async def get_runtime_web_search_config(db: AsyncSession) -> WebSearchConfig:
runtime_record = await get_setting_record(db, "external_integrations")
payload = merge_with_defaults(
"external_integrations",
runtime_record.payload if runtime_record else None,
)
return _runtime_config_from_web_search_payload(payload.get("web_search") or {})
async def get_web_search_client(db: AsyncSession) -> WebSearchClient:
return WebSearchClient(await get_runtime_web_search_config(db))
async def get_runtime_ai_provider_config(db: AsyncSession) -> dict:
runtime_record = await get_setting_record(db, "external_integrations")
payload = merge_with_defaults(
@@ -459,6 +683,7 @@ async def serialize_external_integrations(db: AsyncSession) -> dict:
runtime_setting.payload if runtime_setting else None,
)
normalized_ai = _normalize_ai_provider_payload(raw_payload.get("ai_provider") or {})
normalized_web_search = _normalize_web_search_payload(raw_payload.get("web_search") or {})
default_provider = normalized_ai["default_provider"]
providers_payload: dict[str, dict] = {}
for provider in sorted({
@@ -480,6 +705,32 @@ async def serialize_external_integrations(db: AsyncSession) -> dict:
"source": "runtime" if provider_config.get("api_key") else (api_key_source or "preset"),
}
display_llm_config = providers_payload.get(default_provider) or _provider_defaults(default_provider)
web_search_providers_payload: dict[str, dict] = {}
for provider in sorted({
*[preset["provider"] for preset in list_web_search_provider_presets()],
*normalized_web_search["providers"].keys(),
normalized_web_search["default_provider"],
}):
provider_id = normalize_web_search_provider(provider)
provider_config = (
normalized_web_search["providers"].get(provider_id)
or _web_search_provider_defaults(provider_id)
)
api_key, api_key_source = _resolve_web_search_api_key(provider_id, provider_config)
web_search_providers_payload[provider_id] = {
**{
key: value
for key, value in provider_config.items()
if key != "api_key"
},
"provider": provider_id,
"api_key": _mask_secret(api_key, api_key_source),
"source": "runtime" if provider_config.get("api_key") else (api_key_source or "preset"),
}
display_web_search_config = (
web_search_providers_payload.get(normalized_web_search["default_provider"])
or _web_search_provider_defaults(normalized_web_search["default_provider"])
)
barentswatch_record = await get_barentswatch_config_record(db)
barentswatch_auth = barentswatch_record.auth_config if barentswatch_record else {}
barentswatch_auth = barentswatch_auth or {}
@@ -509,6 +760,28 @@ async def serialize_external_integrations(db: AsyncSession) -> dict:
),
"source": resolved_barentswatch.credential_source,
},
"web_search": {
"enabled": normalized_web_search["enabled"],
"default_provider": normalized_web_search["default_provider"],
"provider": normalized_web_search["default_provider"],
"base_url": display_web_search_config.get("base_url") or "",
"api_key": display_web_search_config.get("api_key") or _mask_secret(None),
"providers": web_search_providers_payload,
"max_results": int(display_web_search_config.get("max_results") or 5),
"timeout_seconds": int(display_web_search_config.get("timeout_seconds") or 20),
"endpoint_path": display_web_search_config.get("endpoint_path") or "",
"search_depth": display_web_search_config.get("search_depth") or "basic",
"engine": display_web_search_config.get("engine") or "google",
"include_answer": bool(display_web_search_config.get("include_answer", False)),
"include_raw_content": bool(display_web_search_config.get("include_raw_content", False)),
"include_text": bool(display_web_search_config.get("include_text", False)),
"categories": display_web_search_config.get("categories") or "general",
"engines": display_web_search_config.get("engines") or [],
"search_path": display_web_search_config.get("search_path") or "",
"scrape_path": display_web_search_config.get("scrape_path") or "",
"scrape_formats": display_web_search_config.get("scrape_formats") or ["markdown"],
"source": "runtime" if runtime_setting else "env",
},
}
@@ -518,8 +791,13 @@ async def save_external_integrations_payload(
) -> dict:
current_payload = await get_setting_payload(db, "external_integrations")
ai_payload = _build_ai_provider_payload(current_payload, update.ai_provider)
web_search_payload = _build_web_search_payload(current_payload, update.web_search)
await save_setting_payload(db, "external_integrations", {"ai_provider": ai_payload})
await save_setting_payload(
db,
"external_integrations",
{"ai_provider": ai_payload, "web_search": web_search_payload},
)
default_endpoint = get_data_sources_config().get_yaml_url("barentswatch_vessels")
barentswatch_record = await get_barentswatch_config_record(db)
@@ -754,7 +1032,12 @@ async def connect_ai_provider_integration(
constraints=["回复尽量简短。"],
)
)
await save_setting_payload(db, "external_integrations", {"ai_provider": draft_ai_payload})
current_web_search = _normalize_web_search_payload(current_payload.get("web_search") or {})
await save_setting_payload(
db,
"external_integrations",
{"ai_provider": draft_ai_payload, "web_search": current_web_search},
)
return {
"success": True,
"connected": True,
@@ -799,6 +1082,74 @@ async def reveal_ai_provider_secrets(
}
@router.get("/integrations/web-search/presets")
async def get_web_search_presets(
current_user: User = Depends(get_current_user),
):
return {"data": list_web_search_provider_presets()}
@router.get("/integrations/web-search/secrets")
async def reveal_web_search_secrets(
provider: str = Query(default=""),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
current_payload = await get_setting_payload(db, "external_integrations")
web_search_payload = _normalize_web_search_payload(current_payload.get("web_search") or {})
provider_id = normalize_web_search_provider(provider or web_search_payload["default_provider"])
provider_config = (
web_search_payload["providers"].get(provider_id)
or _web_search_provider_defaults(provider_id)
)
api_key, api_key_source = _resolve_web_search_api_key(provider_id, provider_config)
return {
"provider": provider_id,
"api_key": api_key,
"api_key_source": api_key_source,
}
@router.post("/integrations/web-search/connect")
async def connect_web_search_integration(
payload: WebSearchIntegrationUpdate,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
current_payload = await get_setting_payload(db, "external_integrations")
draft_web_search_payload = _build_web_search_payload(current_payload, payload)
runtime_config = _runtime_config_from_web_search_payload(draft_web_search_payload)
client = WebSearchClient(runtime_config)
try:
results = await client.test_connection()
return {
"success": True,
"connected": True,
"message": "WebSearch 连接成功。",
"provider": runtime_config.default_provider,
"results": [item.model_dump(mode="json") for item in results[:3]],
}
except WebSearchConfigurationError as exc:
return {
"success": False,
"connected": False,
"message": str(exc),
}
except WebSearchError as exc:
return {
"success": False,
"connected": False,
"message": str(exc),
}
except Exception as exc:
return {
"success": False,
"connected": False,
"message": f"WebSearch 连接测试失败: {exc}",
}
@router.get("/credential-guides/{provider}")
async def read_credential_guide(
provider: str,
@@ -819,7 +1170,15 @@ async def generate_provider_credential_guide(
ai_client: AIProviderClient = Depends(get_ai_provider_client),
):
try:
return {"guide": await generate_credential_guide(db, provider, ai_client)}
web_search_client = await get_web_search_client(db)
return {
"guide": await generate_credential_guide(
db,
provider,
ai_client,
web_search_client,
)
}
except ValueError as exc:
raise HTTPException(status_code=404, detail=str(exc)) from exc

View File

@@ -38,8 +38,12 @@ from app.services.compute_center_locations import (
upsert_compute_center_location,
)
from app.services.ai_client import get_ai_provider_client
from app.api.v1.settings import get_web_search_client
from app.services.collectors.bgp_common import RIPE_RIS_COLLECTOR_COORDS
from app.services.location.llm_fallback import collect_llm_location_fallback_candidate
from app.services.location.llm_fallback import (
collect_llm_location_fallback_candidate,
collect_location_search_evidence,
)
from app.services.persistent_logs import record_system_log
from app.services.vessel_ais_aggregation import (
build_field_conflict_candidates,
@@ -1879,21 +1883,33 @@ async def collect_compute_center_location(
city=city,
country=country,
)
llm_result = None
try:
web_search_client = await get_web_search_client(db)
search_result = await collect_location_search_evidence(
web_search_client=web_search_client,
query=query,
entity_type="compute_center",
)
attempted_queries = [*attempted_queries, *search_result.attempted_queries]
if not search_result.evidence:
llm_failure_reason = search_result.failure_reason
raise RuntimeError(search_result.failure_reason or "no WebSearch evidence")
provider_client = await get_ai_provider_client(db)
llm_result = await collect_llm_location_fallback_candidate(
provider_client=provider_client,
query=query,
entity_type="compute_center",
attempted_queries=attempted_queries,
search_evidence=search_result.evidence,
)
except Exception as exc:
llm_result = None
llm_failure_reason = f"LLM location factcheck unavailable: {exc}"
attempted_queries = [
*attempted_queries,
f"llm_factcheck:compute_center:{name or source_id or 'unknown'}",
]
if llm_failure_reason is None:
llm_failure_reason = f"LLM location factcheck unavailable: {exc}"
attempted_queries = [
*attempted_queries,
f"llm_factcheck:compute_center:{name or source_id or 'unknown'}",
]
if llm_result is not None:
attempted_queries = [*attempted_queries, *llm_result.attempted_queries]
candidates = llm_result.candidates

View File

@@ -0,0 +1,7 @@
"""Backend-owned AI tool services.
The services in this package are business tools used by Planet's backend
orchestrators. They intentionally live outside ``aiprovider`` so model transport
stays separate from evidence collection and domain policy.
"""

View File

@@ -0,0 +1,48 @@
from __future__ import annotations
import hashlib
from typing import Iterable
from app.services.ai_tools.schemas import FetchedEvidence, SearchEvidence
def evidence_content_hash(text: str) -> str:
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def normalize_search_evidence(items: Iterable[SearchEvidence], *, limit: int = 5) -> list[dict]:
normalized: list[dict] = []
seen_urls: set[str] = set()
for item in items:
if not item.url or item.url in seen_urls:
continue
seen_urls.add(item.url)
normalized.append(
{
"title": item.title,
"url": item.url,
"snippet": item.snippet,
"content": item.compact_text(),
"score": item.score,
"source_provider": item.source_provider,
"retrieved_at": item.retrieved_at.isoformat(),
"metadata": item.metadata,
}
)
if len(normalized) >= limit:
break
return normalized
def normalize_fetched_evidence(item: FetchedEvidence, *, text_limit: int = 1200) -> dict:
text = " ".join(item.text.split())[:text_limit]
return {
"title": item.title,
"url": item.final_url or item.url,
"text": text,
"content_hash": item.content_hash or evidence_content_hash(item.text),
"extractor": item.extractor,
"retrieved_at": item.retrieved_at.isoformat(),
"metadata": item.metadata,
}

View File

@@ -0,0 +1,63 @@
from __future__ import annotations
from datetime import UTC, datetime
from typing import Any
from pydantic import BaseModel, Field
class SearchEvidence(BaseModel):
title: str = ""
url: str = ""
snippet: str = ""
content: str = ""
score: float | None = None
source_provider: str = ""
retrieved_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
metadata: dict[str, Any] = Field(default_factory=dict)
def compact_text(self, limit: int = 700) -> str:
text = " ".join((self.content or self.snippet or "").split())
return text[:limit]
class FetchedEvidence(BaseModel):
url: str
final_url: str = ""
title: str = ""
text: str = ""
content_hash: str = ""
extractor: str = "basic_html"
retrieved_at: datetime = Field(default_factory=lambda: datetime.now(UTC))
metadata: dict[str, Any] = Field(default_factory=dict)
class WebSearchProviderConfig(BaseModel):
provider: str = "tavily"
base_url: str = ""
api_key: str = ""
max_results: int = Field(default=5, ge=1, le=20)
timeout_seconds: int = Field(default=20, ge=3, le=120)
endpoint_path: str = ""
search_depth: str = "basic"
engine: str = "google"
include_answer: bool = False
include_raw_content: bool = False
include_text: bool = False
categories: str = "general"
engines: list[str] = Field(default_factory=list)
search_path: str = ""
scrape_path: str = ""
scrape_formats: list[str] = Field(default_factory=lambda: ["markdown"])
class WebSearchConfig(BaseModel):
enabled: bool = False
default_provider: str = "tavily"
provider: str = "tavily"
providers: dict[str, WebSearchProviderConfig] = Field(default_factory=dict)
@property
def active_provider_config(self) -> WebSearchProviderConfig:
return self.providers.get(self.default_provider) or self.providers.get(self.provider) or WebSearchProviderConfig(provider=self.default_provider or self.provider)

View File

@@ -0,0 +1,56 @@
from __future__ import annotations
import hashlib
import httpx
from bs4 import BeautifulSoup
from app.services.ai_tools.schemas import FetchedEvidence
class WebFetchError(RuntimeError):
pass
def _extract_title_and_text(html: str) -> tuple[str, str]:
soup = BeautifulSoup(html, "html.parser")
for tag in soup(["script", "style", "noscript", "svg"]):
tag.decompose()
title = soup.title.get_text(" ", strip=True) if soup.title else ""
main = soup.find("main") or soup.find("article") or soup.body or soup
text = main.get_text("\n", strip=True)
lines = [line.strip() for line in text.splitlines() if line.strip()]
return title, "\n".join(lines)
async def fetch_url_evidence(
url: str,
*,
timeout_seconds: int = 20,
max_bytes: int = 1_500_000,
) -> FetchedEvidence:
if not url:
raise WebFetchError("url is required")
try:
async with httpx.AsyncClient(
timeout=timeout_seconds,
follow_redirects=True,
headers={"User-Agent": "PlanetEvidenceFetcher/1.0"},
) as client:
response = await client.get(url)
response.raise_for_status()
content = response.content[:max_bytes]
except httpx.HTTPError as exc:
raise WebFetchError(f"failed to fetch page: {exc}") from exc
title, text = _extract_title_and_text(content.decode(response.encoding or "utf-8", errors="ignore"))
content_hash = hashlib.sha256(text.encode("utf-8")).hexdigest()
return FetchedEvidence(
url=url,
final_url=str(response.url),
title=title,
text=text,
content_hash=content_hash,
extractor="beautifulsoup_basic",
)

View File

@@ -0,0 +1,391 @@
from __future__ import annotations
from copy import deepcopy
from typing import Any
import httpx
from app.services.ai_tools.schemas import SearchEvidence, WebSearchConfig, WebSearchProviderConfig
WEB_SEARCH_PROVIDER_PRESETS: dict[str, dict[str, Any]] = {
"tavily": {
"provider": "tavily",
"label": "Tavily",
"api_key_env": "TAVILY_API_KEY",
"base_url": "https://api.tavily.com",
"endpoint_path": "/search",
"max_results": 5,
"timeout_seconds": 20,
"search_depth": "basic",
"include_answer": False,
"include_raw_content": False,
},
"brave": {
"provider": "brave",
"label": "Brave Search API",
"api_key_env": "BRAVE_SEARCH_API_KEY",
"base_url": "https://api.search.brave.com",
"endpoint_path": "/res/v1/web/search",
"max_results": 5,
"timeout_seconds": 20,
},
"serpapi": {
"provider": "serpapi",
"label": "SerpAPI",
"api_key_env": "SERPAPI_API_KEY",
"base_url": "https://serpapi.com",
"endpoint_path": "/search.json",
"engine": "google",
"max_results": 5,
"timeout_seconds": 20,
},
"exa": {
"provider": "exa",
"label": "Exa",
"api_key_env": "EXA_API_KEY",
"base_url": "https://api.exa.ai",
"endpoint_path": "/search",
"max_results": 5,
"timeout_seconds": 20,
"include_text": False,
},
"firecrawl": {
"provider": "firecrawl",
"label": "Firecrawl Search / Scrape",
"api_key_env": "FIRECRAWL_API_KEY",
"base_url": "https://api.firecrawl.dev",
"search_path": "/v2/search",
"scrape_path": "/v2/scrape",
"max_results": 5,
"timeout_seconds": 30,
"scrape_formats": ["markdown"],
},
"searxng": {
"provider": "searxng",
"label": "SearXNG",
"api_key_env": "SEARXNG_API_KEY",
"base_url": "http://localhost:8080",
"endpoint_path": "/",
"max_results": 5,
"timeout_seconds": 20,
"categories": "general",
"engines": [],
},
}
class WebSearchError(RuntimeError):
pass
class WebSearchConfigurationError(WebSearchError):
pass
def normalize_web_search_provider(provider: str | None) -> str:
return (provider or "tavily").strip().lower() or "tavily"
def get_web_search_provider_preset(provider: str) -> dict[str, Any]:
provider_id = normalize_web_search_provider(provider)
preset = WEB_SEARCH_PROVIDER_PRESETS.get(provider_id)
if not preset:
raise ValueError(f"Unsupported web search provider: {provider}")
return deepcopy(preset)
def list_web_search_provider_presets() -> list[dict[str, Any]]:
return [get_web_search_provider_preset(provider) for provider in WEB_SEARCH_PROVIDER_PRESETS]
def provider_defaults(provider: str) -> WebSearchProviderConfig:
preset = get_web_search_provider_preset(provider)
return WebSearchProviderConfig(**{
key: value
for key, value in preset.items()
if key in WebSearchProviderConfig.model_fields
})
class WebSearchClient:
def __init__(self, config: WebSearchConfig) -> None:
self.config = config
async def search(
self,
query: str,
*,
max_results: int | None = None,
domains: list[str] | None = None,
freshness_days: int | None = None,
) -> list[SearchEvidence]:
if not self.config.enabled:
raise WebSearchConfigurationError("WebSearch is disabled.")
provider_config = self.config.active_provider_config
provider = normalize_web_search_provider(provider_config.provider)
if provider != "searxng" and not provider_config.api_key:
raise WebSearchConfigurationError(f"{provider} API key is not configured.")
query = " ".join(str(query or "").split())
if not query:
raise WebSearchConfigurationError("search query is required.")
limit = max_results or provider_config.max_results
if provider == "tavily":
return await self._search_tavily(provider_config, query, limit, domains, freshness_days)
if provider == "brave":
return await self._search_brave(provider_config, query, limit, domains)
if provider == "serpapi":
return await self._search_serpapi(provider_config, query, limit)
if provider == "exa":
return await self._search_exa(provider_config, query, limit, domains)
if provider == "firecrawl":
return await self._search_firecrawl(provider_config, query, limit)
if provider == "searxng":
return await self._search_searxng(provider_config, query, limit, domains)
raise WebSearchConfigurationError(f"Unsupported web search provider: {provider}")
async def test_connection(self) -> list[SearchEvidence]:
return await self.search("Planet WebSearch connectivity test", max_results=1)
async def _request_json(
self,
method: str,
url: str,
*,
provider_config: WebSearchProviderConfig,
headers: dict[str, str] | None = None,
params: dict[str, Any] | None = None,
json: dict[str, Any] | None = None,
) -> dict[str, Any]:
try:
async with httpx.AsyncClient(timeout=provider_config.timeout_seconds) as client:
response = await client.request(
method,
url,
headers=headers,
params=params,
json=json,
)
response.raise_for_status()
data = response.json()
except httpx.HTTPStatusError as exc:
detail = exc.response.text or exc.response.reason_phrase
raise WebSearchError(f"{provider_config.provider} request failed: {detail}") from exc
except httpx.HTTPError as exc:
raise WebSearchError(f"{provider_config.provider} request failed: {exc}") from exc
except ValueError as exc:
raise WebSearchError(f"{provider_config.provider} returned invalid JSON") from exc
return data if isinstance(data, dict) else {}
async def _search_tavily(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
domains: list[str] | None,
freshness_days: int | None,
) -> list[SearchEvidence]:
body: dict[str, Any] = {
"api_key": config.api_key,
"query": query,
"max_results": max_results,
"search_depth": config.search_depth or "basic",
"include_answer": config.include_answer,
"include_raw_content": config.include_raw_content,
}
if domains:
body["include_domains"] = domains
if freshness_days:
body["days"] = freshness_days
data = await self._request_json(
"POST",
_join_url(config.base_url, config.endpoint_path or "/search"),
provider_config=config,
json=body,
)
return [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("url") or ""),
snippet=str(item.get("content") or ""),
content=str(item.get("raw_content") or ""),
score=_float_or_none(item.get("score")),
source_provider="tavily",
metadata={"query": data.get("query") or query},
)
for item in data.get("results") or []
if isinstance(item, dict) and item.get("url")
]
async def _search_brave(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
domains: list[str] | None,
) -> list[SearchEvidence]:
search_query = query
if domains:
search_query = f"{query} " + " ".join(f"site:{domain}" for domain in domains)
data = await self._request_json(
"GET",
_join_url(config.base_url, config.endpoint_path or "/res/v1/web/search"),
provider_config=config,
headers={"X-Subscription-Token": config.api_key},
params={"q": search_query, "count": max_results},
)
results = (data.get("web") or {}).get("results") or []
return [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("url") or ""),
snippet=str(item.get("description") or ""),
source_provider="brave",
metadata={"age": item.get("age")},
)
for item in results
if isinstance(item, dict) and item.get("url")
]
async def _search_serpapi(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
) -> list[SearchEvidence]:
data = await self._request_json(
"GET",
_join_url(config.base_url, config.endpoint_path or "/search.json"),
provider_config=config,
params={
"api_key": config.api_key,
"engine": config.engine or "google",
"q": query,
"num": max_results,
},
)
return [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("link") or ""),
snippet=str(item.get("snippet") or ""),
source_provider="serpapi",
metadata={"position": item.get("position")},
)
for item in data.get("organic_results") or []
if isinstance(item, dict) and item.get("link")
]
async def _search_exa(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
domains: list[str] | None,
) -> list[SearchEvidence]:
body: dict[str, Any] = {
"query": query,
"numResults": max_results,
}
if domains:
body["includeDomains"] = domains
if config.include_text:
body["contents"] = {"text": True}
data = await self._request_json(
"POST",
_join_url(config.base_url, config.endpoint_path or "/search"),
provider_config=config,
headers={"Authorization": f"Bearer {config.api_key}"},
json=body,
)
return [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("url") or ""),
snippet=str(item.get("summary") or ""),
content=str(item.get("text") or ""),
score=_float_or_none(item.get("score")),
source_provider="exa",
metadata={"id": item.get("id")},
)
for item in data.get("results") or []
if isinstance(item, dict) and item.get("url")
]
async def _search_firecrawl(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
) -> list[SearchEvidence]:
data = await self._request_json(
"POST",
_join_url(config.base_url, config.search_path or "/v2/search"),
provider_config=config,
headers={"Authorization": f"Bearer {config.api_key}"},
json={"query": query, "limit": max_results},
)
raw_results = data.get("data") or data.get("results") or []
return [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("url") or item.get("sourceURL") or ""),
snippet=str(item.get("description") or item.get("markdown") or ""),
source_provider="firecrawl",
metadata={"status": item.get("status")},
)
for item in raw_results
if isinstance(item, dict) and (item.get("url") or item.get("sourceURL"))
]
async def _search_searxng(
self,
config: WebSearchProviderConfig,
query: str,
max_results: int,
domains: list[str] | None,
) -> list[SearchEvidence]:
search_query = query
if domains:
search_query = f"{query} " + " ".join(f"site:{domain}" for domain in domains)
params: dict[str, Any] = {
"q": search_query,
"format": "json",
"categories": config.categories or "general",
}
if config.engines:
params["engines"] = ",".join(config.engines)
headers = {"Authorization": f"Bearer {config.api_key}"} if config.api_key else None
data = await self._request_json(
"GET",
_join_url(config.base_url, config.endpoint_path or "/"),
provider_config=config,
headers=headers,
params=params,
)
results = data.get("results") or []
evidence = [
SearchEvidence(
title=str(item.get("title") or ""),
url=str(item.get("url") or ""),
snippet=str(item.get("content") or ""),
score=_float_or_none(item.get("score")),
source_provider="searxng",
metadata={"engine": item.get("engine")},
)
for item in results
if isinstance(item, dict) and item.get("url")
]
return evidence[:max_results]
def _join_url(base_url: str, path: str) -> str:
return f"{(base_url or '').rstrip('/')}/{(path or '').lstrip('/')}"
def _float_or_none(value: Any) -> float | None:
try:
return float(value)
except (TypeError, ValueError):
return None

View File

@@ -10,6 +10,8 @@ from sqlalchemy import select
from app.models.system_setting import SystemSetting
from app.schemas.ai import SituationalAnalysisRequest
from app.services.ai_client import AIProviderClient
from app.services.ai_tools.evidence_store import normalize_search_evidence
from app.services.ai_tools.web_search import WebSearchClient, WebSearchError
CREDENTIAL_GUIDES_CATEGORY = "collector_credential_guides"
@@ -153,10 +155,26 @@ async def get_credential_guide(db, provider: str) -> dict[str, Any]:
"markdown": custom.get("markdown") if custom else default.markdown,
"prompt": default.prompt,
"source": "ai" if custom else "default",
"sources": custom.get("sources", []) if custom else [],
"verification_status": (
custom.get("verification_status", "verified_with_search_evidence")
if custom
else "default_unverified"
),
"verification_error": custom.get("verification_error") if custom else None,
}
async def save_credential_guide(db, provider: str, title: str, markdown: str) -> dict[str, Any]:
async def save_credential_guide(
db,
provider: str,
title: str,
markdown: str,
*,
sources: list[dict[str, Any]] | None = None,
verification_status: str = "verified_with_search_evidence",
verification_error: str | None = None,
) -> dict[str, Any]:
default = DEFAULT_CREDENTIAL_GUIDES.get(provider)
if default is None:
raise ValueError(f"Unsupported credential guide provider: {provider}")
@@ -165,6 +183,9 @@ async def save_credential_guide(db, provider: str, title: str, markdown: str) ->
store[provider] = {
"title": title or default.title,
"markdown": markdown,
"sources": sources or [],
"verification_status": verification_status,
"verification_error": verification_error,
}
if record is None:
db.add(SystemSetting(category=CREDENTIAL_GUIDES_CATEGORY, payload=store))
@@ -192,28 +213,58 @@ async def generate_credential_guide(
db,
provider: str,
ai_client: AIProviderClient,
web_search_client: WebSearchClient | None = None,
) -> dict[str, Any]:
default = DEFAULT_CREDENTIAL_GUIDES.get(provider)
if default is None:
raise ValueError(f"Unsupported credential guide provider: {provider}")
search_evidence: list[dict[str, Any]] = []
search_error: str | None = None
if web_search_client is not None:
try:
evidence = await web_search_client.search(
_credential_guide_search_query(default),
max_results=5,
)
search_evidence = normalize_search_evidence(evidence, limit=5)
except WebSearchError as exc:
search_error = str(exc)
except Exception as exc:
search_error = f"WebSearch unavailable: {exc}"
if not search_evidence:
guide = await get_credential_guide(db, provider)
guide["verification_status"] = "unverified_no_search_evidence"
guide["verification_error"] = search_error
guide["sources"] = []
return guide
response = await ai_client.analyze(
SituationalAnalysisRequest(
title=f"Generate credential guide for {provider}",
objective=default.prompt,
objective=(
default.prompt
+ "\n只能根据 context.search_evidence 中的来源生成教程;"
+ "如果证据不足,明确说明需要以官方页面为准。"
),
context={
"provider": provider,
"current_default_guide": default.markdown,
"product_context": "Planet collector credential settings",
"search_evidence": search_evidence,
},
observations=[
"Use concise Chinese markdown.",
"Prefer stable concepts over brittle UI labels.",
"Include verification and troubleshooting steps.",
"Include a short sources section with the provided URLs.",
],
constraints=[
"Do not ask the user for secrets.",
"Do not include fabricated screenshots.",
"Do not invent source URLs or product UI labels.",
"Use only the provided search_evidence as factual support.",
"Return markdown only.",
],
)
@@ -221,4 +272,19 @@ async def generate_credential_guide(
markdown = response.content.strip()
if not markdown:
markdown = default.markdown
return await save_credential_guide(db, provider, default.title, markdown)
return await save_credential_guide(
db,
provider,
default.title,
markdown,
sources=search_evidence,
verification_status="verified_with_search_evidence",
)
def _credential_guide_search_query(default: CredentialGuideDefault) -> str:
if default.provider == "barentswatch":
return "BarentsWatch developer tutorial AIS API OAuth client credentials"
if default.provider == "aisstream":
return "AISStream API key documentation websocket stream"
return f"{default.provider} API credentials documentation"

View File

@@ -10,6 +10,8 @@ from typing import Any, Iterable
from app.core.countries import COUNTRY_ENTRIES, normalize_country
from app.schemas.ai import SituationalAnalysisRequest
from app.services.ai_client import AIProviderClient
from app.services.ai_tools.evidence_store import normalize_search_evidence
from app.services.ai_tools.web_search import WebSearchClient, WebSearchError
from app.services.location.models import LocationCandidate, LocationQuery
from app.services.location.resolvers.nominatim import build_default_nominatim_geocoder
from app.services.location.text import (
@@ -69,6 +71,13 @@ class LocationLLMFallbackResult:
failure_reason: str | None = None
@dataclass(frozen=True)
class LocationSearchEvidenceResult:
evidence: list[dict[str, Any]]
attempted_queries: list[str]
failure_reason: str | None = None
@dataclass(frozen=True)
class LocationEvidenceScore:
score: float
@@ -751,6 +760,10 @@ def _candidate_from_payload(
"name_location_hint": evidence_score.name_location_hint,
},
},
raw_payload={
"llm_payload": payload,
"search_evidence": payload.get("search_evidence") or [],
},
), None
@@ -802,6 +815,61 @@ def _observations(query: LocationQuery, attempted_queries: Iterable[str]) -> lis
return observations
def _location_search_query(query: LocationQuery, entity_type: str) -> str:
extra = query.extra or {}
parts = [
coerce_str(query.name),
coerce_str(extra.get("site")),
coerce_str(extra.get("operator")),
coerce_str(extra.get("organization")),
coerce_str(query.city),
coerce_str(query.country),
"physical location",
]
if entity_type == "bgp_collector":
parts.append("route collector city")
elif entity_type == "compute_center":
parts.append("datacenter supercomputer facility city")
return " ".join(part for part in parts if part)
async def collect_location_search_evidence(
*,
web_search_client: WebSearchClient,
query: LocationQuery,
entity_type: str,
max_results: int = 5,
) -> LocationSearchEvidenceResult:
search_query = _location_search_query(query, entity_type)
attempt = f"web_search:{entity_type}:{search_query}"
try:
evidence = await web_search_client.search(search_query, max_results=max_results)
except WebSearchError as exc:
return LocationSearchEvidenceResult(
evidence=[],
attempted_queries=[attempt],
failure_reason=f"WebSearch location evidence failed: {exc}",
)
except Exception as exc:
return LocationSearchEvidenceResult(
evidence=[],
attempted_queries=[attempt],
failure_reason=f"WebSearch location evidence unavailable: {exc}",
)
normalized = normalize_search_evidence(evidence, limit=max_results)
if not normalized:
return LocationSearchEvidenceResult(
evidence=[],
attempted_queries=[attempt],
failure_reason="WebSearch returned no usable location evidence.",
)
return LocationSearchEvidenceResult(
evidence=normalized,
attempted_queries=[attempt],
failure_reason=None,
)
async def _repair_location_payload_from_text(
*,
provider_client: AIProviderClient,
@@ -862,6 +930,7 @@ async def collect_llm_location_fallback_candidate(
query: LocationQuery,
entity_type: str,
attempted_queries: Iterable[str] = (),
search_evidence: list[dict[str, Any]] | None = None,
min_confidence: float = DEFAULT_MIN_CONFIDENCE,
) -> LocationLLMFallbackResult:
"""Ask the configured LLM for one fact-checked location candidate.
@@ -871,6 +940,12 @@ async def collect_llm_location_fallback_candidate(
use this in user-triggered collection flows.
"""
attempt = f"llm_factcheck:{entity_type}:{coerce_str(query.name) or 'unknown'}"
if search_evidence is not None and not search_evidence:
return LocationLLMFallbackResult(
candidates=[],
attempted_queries=[attempt],
failure_reason="LLM location factcheck skipped: no WebSearch evidence.",
)
request = SituationalAnalysisRequest(
title=f"Location factcheck fallback for {entity_type}",
objective=(
@@ -881,6 +956,7 @@ async def collect_llm_location_fallback_candidate(
context={
"entity_type": entity_type,
"location_query": _query_context(query),
"search_evidence": search_evidence or [],
"required_json_schema": {
"latitude": "number",
"longitude": "number",
@@ -905,6 +981,7 @@ async def collect_llm_location_fallback_candidate(
"Calibrate model confidence using this rubric: 0.85-1.0 for exact facility coordinates backed by an authoritative source; 0.70-0.84 for a confirmed facility/campus with strong public evidence; 0.55-0.69 for a confirmed city-level location backed by credible sources but without exact facility coordinates; 0.35-0.54 for weak or ambiguous city evidence; below 0.35 when the location is mostly a guess.",
"Return evidence as objects when possible, including source, url, source_type, and entity_match.",
"Include source names or URLs in evidence when known. The backend will recompute the final confidence from model confidence plus evidence quality.",
"If search_evidence is provided, use only that evidence as factual support.",
"Prefer the facility/site if known; otherwise use the best supported city.",
],
)
@@ -939,6 +1016,23 @@ async def collect_llm_location_fallback_candidate(
),
)
payload = _normalize_llm_payload(payload)
if search_evidence:
payload["search_evidence"] = search_evidence
existing_evidence = _evidence_items(payload.get("evidence"))
payload["evidence"] = [
*existing_evidence,
*[
{
"source": item.get("title") or item.get("url"),
"url": item.get("url"),
"text": item.get("snippet") or item.get("content"),
"source_type": "web_search",
"entity_match": True,
}
for item in search_evidence
if isinstance(item, dict)
],
]
latitude, longitude = _extract_llm_coordinates(payload)
city_geocode_failure = None
if latitude in (None, 0.0) or longitude in (None, 0.0):

View File

@@ -51,6 +51,7 @@ class LocationCandidate:
matched_location_name: str | None = None
location_verified_at: str | None = None
suggested_registry_entry: dict[str, Any] | None = None
raw_payload: dict[str, Any] | None = None
def to_dict(self) -> dict[str, Any]:
return {
@@ -70,6 +71,7 @@ class LocationCandidate:
"matched_location_name": self.matched_location_name,
"location_verified_at": self.location_verified_at,
"suggested_registry_entry": self.suggested_registry_entry,
"raw_payload": self.raw_payload,
}

View File

@@ -137,6 +137,20 @@ async def test_collect_bgp_collector_location_uses_llm_when_candidates_empty(mon
lambda **_kwargs: ([], ["Lyon, France"]),
)
from app.services.location.llm_fallback import LocationSearchEvidenceResult
async def _search_evidence(**_kwargs):
return LocationSearchEvidenceResult(
evidence=[
{
"title": "RRC source",
"url": "https://example.test/rrc",
"snippet": "rrc-mystery is in Lyon.",
}
],
attempted_queries=["web_search:bgp_collector:rrc-mystery Lyon France physical location route collector city"],
)
async def _fallback(**_kwargs):
return LocationLLMFallbackResult(
candidates=[llm_candidate],
@@ -144,6 +158,8 @@ async def test_collect_bgp_collector_location_uses_llm_when_candidates_empty(mon
)
monkeypatch.setattr(bgp_api, "get_ai_provider_client", AsyncMock(return_value=object()))
monkeypatch.setattr(bgp_api, "get_web_search_client", AsyncMock(return_value=object()))
monkeypatch.setattr(bgp_api, "collect_location_search_evidence", _search_evidence)
monkeypatch.setattr(bgp_api, "collect_llm_location_fallback_candidate", _fallback)
response = await bgp_api.collect_bgp_collector_location(
@@ -158,6 +174,7 @@ async def test_collect_bgp_collector_location_uses_llm_when_candidates_empty(mon
assert response["best_candidate"]["needs_confirmation"] is True
assert response["attempted_queries"] == [
"Lyon, France",
"web_search:bgp_collector:rrc-mystery Lyon France physical location route collector city",
"llm_factcheck:bgp_collector:rrc-mystery",
]

View File

@@ -569,7 +569,19 @@ async def test_collect_compute_center_location_uses_llm_when_candidates_empty(mo
lambda **_kwargs: ([], ["Mystery Cluster, France"]),
)
from app.services.location.llm_fallback import LocationLLMFallbackResult
from app.services.location.llm_fallback import LocationLLMFallbackResult, LocationSearchEvidenceResult
async def _search_evidence(**_kwargs):
return LocationSearchEvidenceResult(
evidence=[
{
"title": "Mystery Cluster source",
"url": "https://example.test/mystery",
"snippet": "Mystery Cluster is in Lyon.",
}
],
attempted_queries=["web_search:compute_center:Mystery Cluster France physical location"],
)
async def _fallback(**_kwargs):
return LocationLLMFallbackResult(
@@ -578,6 +590,8 @@ async def test_collect_compute_center_location_uses_llm_when_candidates_empty(mo
)
monkeypatch.setattr(visualization_api, "get_ai_provider_client", AsyncMock(return_value=object()))
monkeypatch.setattr(visualization_api, "get_web_search_client", AsyncMock(return_value=object()))
monkeypatch.setattr(visualization_api, "collect_location_search_evidence", _search_evidence)
monkeypatch.setattr(visualization_api, "collect_llm_location_fallback_candidate", _fallback)
response = await visualization_api.collect_compute_center_location(
@@ -595,6 +609,7 @@ async def test_collect_compute_center_location_uses_llm_when_candidates_empty(mo
assert response["best_candidate"]["needs_confirmation"] is True
assert response["attempted_queries"] == [
"Mystery Cluster, France",
"web_search:compute_center:Mystery Cluster France physical location",
"llm_factcheck:compute_center:Mystery Cluster",
]

View File

@@ -0,0 +1,261 @@
from types import SimpleNamespace
import pytest
from app.api.v1 import settings as settings_api
from app.api.v1.settings import (
WebSearchIntegrationUpdate,
_build_web_search_payload,
_mask_secret,
_normalize_web_search_payload,
_resolve_web_search_api_key,
)
from app.services.ai_tools.schemas import WebSearchConfig, WebSearchProviderConfig
from app.services.ai_tools.web_search import WebSearchClient
from app.services.credential_guides import generate_credential_guide
from app.services.location.llm_fallback import (
collect_llm_location_fallback_candidate,
collect_location_search_evidence,
)
from app.services.location.models import LocationQuery
@pytest.fixture(autouse=True)
def isolated_web_search_env_files(monkeypatch, tmp_path):
env_file = tmp_path / ".env"
monkeypatch.setattr(settings_api, "WEB_SEARCH_ENV_FILES", (env_file,))
return env_file
def test_normalize_web_search_payload_adds_default_provider():
payload = _normalize_web_search_payload({})
assert payload["default_provider"] == "tavily"
assert payload["providers"]["tavily"]["base_url"] == "https://api.tavily.com"
def test_web_search_key_prefers_provider_env(isolated_web_search_env_files):
isolated_web_search_env_files.write_text(
"TAVILY_API_KEY=tavily-env-key\nWEB_SEARCH_API_KEY=generic-search-key\n",
encoding="utf-8",
)
value, source = _resolve_web_search_api_key("tavily", {"api_key": ""})
assert value == "tavily-env-key"
assert source == "env_file"
def test_build_web_search_payload_keeps_saved_key_when_preview_submitted():
current = {
"web_search": {
"default_provider": "tavily",
"providers": {
"tavily": {
"provider": "tavily",
"api_key": "tvly-old-secret",
},
},
}
}
update = WebSearchIntegrationUpdate(
enabled=True,
provider="tavily",
base_url="https://api.tavily.com",
api_key=_mask_secret("tvly-old-secret")["preview"],
)
payload = _build_web_search_payload(current, update)
assert payload["enabled"] is True
assert payload["providers"]["tavily"]["api_key"] == "tvly-old-secret"
@pytest.mark.asyncio
async def test_tavily_adapter_normalizes_results(monkeypatch):
config = WebSearchConfig(
enabled=True,
default_provider="tavily",
provider="tavily",
providers={
"tavily": WebSearchProviderConfig(
provider="tavily",
base_url="https://api.tavily.com",
api_key="key",
)
},
)
client = WebSearchClient(config)
async def fake_request_json(*args, **kwargs):
return {
"query": "Alem.Cloud",
"results": [
{
"title": "Alem.Cloud official",
"url": "https://example.test/alem",
"content": "Alem.Cloud is in Astana.",
"score": 0.9,
}
],
}
monkeypatch.setattr(client, "_request_json", fake_request_json)
results = await client.search("Alem.Cloud")
assert results[0].source_provider == "tavily"
assert results[0].url == "https://example.test/alem"
assert "Astana" in results[0].snippet
@pytest.mark.asyncio
async def test_searxng_adapter_allows_empty_api_key(monkeypatch):
config = WebSearchConfig(
enabled=True,
default_provider="searxng",
provider="searxng",
providers={
"searxng": WebSearchProviderConfig(
provider="searxng",
base_url="http://localhost:8080",
api_key="",
)
},
)
client = WebSearchClient(config)
async def fake_request_json(*args, **kwargs):
return {
"results": [
{
"title": "TAIPEI-1",
"url": "https://example.test/taipei",
"content": "TAIPEI-1 is in Taipei.",
"score": 2,
"engine": "duckduckgo",
}
],
}
monkeypatch.setattr(client, "_request_json", fake_request_json)
results = await client.search("TAIPEI-1")
assert results[0].source_provider == "searxng"
assert results[0].metadata["engine"] == "duckduckgo"
@pytest.mark.asyncio
async def test_location_search_evidence_returns_failure_on_empty_results(monkeypatch):
class EmptySearchClient:
async def search(self, *args, **kwargs):
return []
result = await collect_location_search_evidence(
web_search_client=EmptySearchClient(),
query=LocationQuery(name="TAIPEI-1", country="Taiwan"),
entity_type="compute_center",
)
assert result.evidence == []
assert "no usable" in result.failure_reason
@pytest.mark.asyncio
async def test_llm_location_fallback_skips_when_search_evidence_empty():
class ExplodingAIClient:
async def analyze(self, *_args, **_kwargs):
raise AssertionError("LLM should not be called without evidence")
result = await collect_llm_location_fallback_candidate(
provider_client=ExplodingAIClient(),
query=LocationQuery(name="TAIPEI-1", country="Taiwan"),
entity_type="compute_center",
search_evidence=[],
)
assert result.candidates == []
assert "no WebSearch evidence" in result.failure_reason
@pytest.mark.asyncio
async def test_credential_guide_keeps_default_without_search_evidence():
class EmptySearchClient:
async def search(self, *args, **kwargs):
return []
class ExplodingAIClient:
async def analyze(self, *_args, **_kwargs):
raise AssertionError("AI should not be called without search evidence")
async def fake_get_store(_db):
return None, {}
import app.services.credential_guides as credential_guides
original = credential_guides._get_guide_store
credential_guides._get_guide_store = fake_get_store
try:
guide = await generate_credential_guide(
object(),
"barentswatch",
ExplodingAIClient(),
EmptySearchClient(),
)
finally:
credential_guides._get_guide_store = original
assert guide["source"] == "default"
assert guide["verification_status"] == "unverified_no_search_evidence"
@pytest.mark.asyncio
async def test_credential_guide_uses_search_evidence(monkeypatch):
class SearchClient:
async def search(self, *args, **kwargs):
from app.services.ai_tools.schemas import SearchEvidence
return [
SearchEvidence(
title="Official docs",
url="https://docs.example.test",
snippet="Create an AIS client.",
source_provider="tavily",
)
]
class AIClient:
async def analyze(self, payload):
assert payload.context["search_evidence"]
return SimpleNamespace(content="## Generated\n\nSources included.")
saved = {}
async def fake_get_store(_db):
return None, saved
async def fake_save(db, provider, title, markdown, **metadata):
return {
"provider": provider,
"title": title,
"markdown": markdown,
"source": "ai",
**metadata,
}
import app.services.credential_guides as credential_guides
monkeypatch.setattr(credential_guides, "_get_guide_store", fake_get_store)
monkeypatch.setattr(credential_guides, "save_credential_guide", fake_save)
guide = await generate_credential_guide(
object(),
"barentswatch",
AIClient(),
SearchClient(),
)
assert guide["source"] == "ai"
assert guide["verification_status"] == "verified_with_search_evidence"
assert guide["sources"][0]["url"] == "https://docs.example.test"