release: bump version to 0.51.0
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
7
backend/app/services/ai_tools/__init__.py
Normal file
7
backend/app/services/ai_tools/__init__.py
Normal 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.
|
||||
"""
|
||||
|
||||
48
backend/app/services/ai_tools/evidence_store.py
Normal file
48
backend/app/services/ai_tools/evidence_store.py
Normal 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,
|
||||
}
|
||||
|
||||
63
backend/app/services/ai_tools/schemas.py
Normal file
63
backend/app/services/ai_tools/schemas.py
Normal 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)
|
||||
|
||||
56
backend/app/services/ai_tools/web_fetch.py
Normal file
56
backend/app/services/ai_tools/web_fetch.py
Normal 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",
|
||||
)
|
||||
|
||||
391
backend/app/services/ai_tools/web_search.py
Normal file
391
backend/app/services/ai_tools/web_search.py
Normal 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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
261
backend/tests/test_web_search_tools.py
Normal file
261
backend/tests/test_web_search_tools.py
Normal 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"
|
||||
Reference in New Issue
Block a user