From ad0190eea5df6e16b47508d7499c5a966c5aa33b Mon Sep 17 00:00:00 2001 From: Eric Hills <53243273+ebhills@users.noreply.github.com> Date: Wed, 30 Sep 2026 13:50:08 -0500 Subject: [PATCH] Configure Gemini retrieval thinking and per-URL deadlines --- docs/ai_configuration.md | 44 ++++-- docs/search_retrieve_link_content.md | 46 +++++- pytest-local.ini | 1 + requirements.txt | 2 +- tests/test_ai_caller_config.py | 191 ++++++++++++++++++++++-- tests/test_ai_config.py | 8 +- tests/test_search_retrieve_templates.py | 98 +++++++++++- wrangles/ai_defaults.yml | 35 ++++- wrangles/clients/gemini.py | 104 +++++++++++-- wrangles/recipe_wrangles/search.py | 26 +++- wrangles/search.py | 37 ++++- 11 files changed, 534 insertions(+), 58 deletions(-) diff --git a/docs/ai_configuration.md b/docs/ai_configuration.md index 4624d034..d525b165 100644 --- a/docs/ai_configuration.md +++ b/docs/ai_configuration.md @@ -90,8 +90,9 @@ operation defaults; operation defaults take precedence. Explicit caller arguments take precedence over resolved defaults, including `False` and `0`. Endpoint overrides remain available in APIs that already expose them. -All packaged operations default to one additional retry after the first attempt. -Set `retries: 0` to disable retries. Temperature is model-specific: modern OpenAI +Packaged operations default to one additional retry after the first attempt, +except Gemini URL retrieval, which defaults to no retries. Set `retries: 0` +to disable retries. Temperature is model-specific: modern OpenAI models and Gemini URL retrieval leave it unset, while legacy GPT-4o Chat Completions keeps `0.2`. Unset temperature uses the provider's default (`1.0` for Gemini 3). @@ -133,13 +134,15 @@ returned dictionary does not alter cached configuration. | `ai.choose`, `ai.score`, `ai.true_false`, `ai.answers` | Python and recipe structured answers | Provider, model, endpoint, concurrency, timeout, retries, cache | | `extract.ai` | Python and recipe extraction | Model, endpoints, model tuning, concurrency, timeout, retries, strictness, storage, cache, prompt | | `embeddings` | `openai.embeddings` and recipe `create.embeddings` | Provider, model, endpoint, batch size, concurrency, timeout, retries, precision, dimensions, Jina task/normalization/truncation | -| `search.retrieve_link_content` | Python, recipe, and Gemini URL-context client | Model, endpoint/API version, concurrency, timeout, retries, temperature/top-p/top-k/token limits/stop sequences | +| `search.retrieve_link_content` | Python, recipe, and Gemini URL-context client | Model, endpoint/API version, concurrency, per-URL deadline, retries, thinking level, temperature/top-p/top-k/token limits/stop sequences | | `generate.ai` | Python and recipe generation | Model, endpoint, reasoning/text tuning, concurrency, timeout, retries, strictness | | `huggingface` | Generic recipe task wrangle | Explicit model, endpoint, timeout, retries, task parameters | -Provider request options use explicit allowlists so runtime settings and catalog -metadata cannot leak into API payloads. Explicit request arguments still override -configured options. Extraction retains `messages`/`examples` aliases and recipe +Configured provider request options use explicit allowlists so runtime settings +and catalog metadata cannot leak into API payloads. Explicit request arguments +still override configured options. Gemini retrieval also accepts caller-supplied +`GenerateContentConfig` options through keyword arguments. Extraction retains +`messages`/`examples` aliases and recipe output-shape controls because existing callers use them. Private transport arguments and unused generation scaffolding have been removed where redundant. @@ -190,15 +193,36 @@ validation, including v3's `separation` value. See the ### Google Gemini -Gemini URL retrieval uses the configured model and Google's URL-context tools. +Gemini URL retrieval defaults to `gemini-3.5-flash` and uses Google's +URL-context tools. Its packaged settings explicitly select +`thinking_level: minimal`, `request_timeout_seconds: 10`, and `retries: 0`. +The Google SDK dependency requires version `1.64.0` or later for these controls. +Use `model_id` to select another model, such as `gemini-3.6-flash`, and +`thinking_level` to override the configured thinking level. The wrapper accepts +`minimal`, `low`, `medium`, and `high`; actual support depends on the model. +Explicit `request_timeout_seconds` overrides the configured deadline. + +The deadline covers one URL's provider request, including all configured retries. +A timeout produces the existing per-URL `Failure` result and preserves output +ordering. Cancellation and client cleanup may take a little longer. Queued URLs +receive their own deadline when their worker starts, so a batch can exceed 10 +seconds when it needs multiple waves of requests. + +Additional caller keyword arguments are forwarded to the SDK's +`GenerateContentConfig` and override configured generation defaults, including +options such as `max_output_tokens`, `top_p`, and `temperature`. Retrieval's +prompt, URL-context tool, response format, and HTTP settings remain managed by +the wrapper. An explicit `thinking_config` replaces the thinking configuration +as a whole, including any selected `thinking_level`. See the +[retrieval guide](search_retrieve_link_content.md) for recipe and Python usage. `search.ai_mode` delegates its underlying model to SerpAPI/Google and has no selectable LLM model in this API. For Google, `endpoints.base_url` is the SDK service root `https://generativelanguage.googleapis.com`. The retrieval operation sets -`api_version: v1beta`; the SDK appends the model and method. With the configured -`gemini-3.8-flash`, the complete request URL is -`https://generativelanguage.googleapis.com/v1beta/models/gemini-3.8-flash:generateContent`. +`api_version: v1beta`; the SDK appends the model and method. With the default +`gemini-3.5-flash`, the complete request URL is +`https://generativelanguage.googleapis.com/v1beta/models/gemini-3.5-flash:generateContent`. Both the base URL and version are configurable. See the [Google API reference](https://ai.google.dev/api/generate-content). Google model names with or without the SDK's `models/` prefix share the same diff --git a/docs/search_retrieve_link_content.md b/docs/search_retrieve_link_content.md index fc2662d4..d0d47aea 100644 --- a/docs/search_retrieve_link_content.md +++ b/docs/search_retrieve_link_content.md @@ -90,8 +90,52 @@ dual output mode: `output: [page_data, page_text]` stores the result list in their existing meanings. See [AI configuration](ai_configuration.md) for model and runtime settings. +## Model, thinking, and request options + +The packaged defaults are `gemini-3.5-flash`, `thinking_level: minimal`, a +10-second deadline per URL, and no retries. Omitted options use the active AI +configuration, so a replacement catalog can change these defaults. Override +`model_id`, `thinking_level`, or `request_timeout_seconds` in the recipe when +needed. For example: + +```yaml +- search.retrieve_link_content: + input: URL + output: Page Content + api_key: ${GEMINI_API_KEY} + model_id: gemini-3.6-flash + thinking_level: minimal + request_timeout_seconds: 10 + max_output_tokens: 2048 + seed: 42 +``` + +`thinking_level` accepts `minimal`, `low`, `medium`, or `high`; support depends +on the selected model. Additional options such as `max_output_tokens`, `top_p`, +and `temperature` are passed as top-level Gemini `GenerateContentConfig` +settings. Explicit options override configured generation defaults. The +installed Google SDK validates these options; their availability depends on +the selected model and SDK version. + +Retrieval manages `system_instruction`, `tools`, `response_mime_type`, +`response_modalities`, and `http_options`; those names and their SDK aliases +cannot be passed as additional options. Use `prompt`, `output_format`, and +`request_timeout_seconds` for the corresponding controls. For advanced thinking +settings, an explicit `thinking_config` replaces the entire thinking +configuration, including any `thinking_level`, rather than merging with it. + +The positive, finite `request_timeout_seconds` value limits the total provider +request time for each URL, including any retries enabled in the AI catalog. +When the deadline expires, that URL returns a `Failure` result with an error +and no extracted content; its position in the output is preserved. Cancellation +and client cleanup may add a little time after the deadline. The deadline +applies separately to each URL after its worker starts, so batches that need +several waves of concurrent requests can take longer than 10 seconds. + ## Direct Python calls Column substitution applies to recipe execution, where a full row is available. Direct calls to `wrangles.search.retrieve_link_content(...)` continue to use -the supplied prompt literally and do not resolve column placeholders. +the supplied prompt literally and do not resolve column placeholders. The +`thinking_level`, `request_timeout_seconds`, and additional Gemini generation +options are also available as Python keyword arguments. diff --git a/pytest-local.ini b/pytest-local.ini index 8fa40a8c..609dd2ad 100644 --- a/pytest-local.ini +++ b/pytest-local.ini @@ -10,6 +10,7 @@ testpaths = tests/test_dataframe.py tests/test_extract_ai_metadata.py tests/test_extract_ai_metadata_context.py + tests/test_lookup_variants.py tests/test_openai_extract_ai.py tests/test_search-ai_mode.py tests/test_search_ai_extraction.py diff --git a/requirements.txt b/requirements.txt index ad0d95ac..070f7347 100644 --- a/requirements.txt +++ b/requirements.txt @@ -50,7 +50,7 @@ pymysql # AI / LLM openai -google-genai>=1.24.0 +google-genai>=1.64.0 # Search serpapi diff --git a/tests/test_ai_caller_config.py b/tests/test_ai_caller_config.py index 7b999c8f..4fa25e7b 100644 --- a/tests/test_ai_caller_config.py +++ b/tests/test_ai_caller_config.py @@ -1,6 +1,7 @@ """Offline contracts for configuration-backed embeddings and URL retrieval.""" import base64 +import asyncio import copy import json from types import SimpleNamespace @@ -63,6 +64,7 @@ def configured_ai(monkeypatch, tmp_path): "default_concurrency": 1, "request_timeout_seconds": 3.25, "temperature": 0, + "retries": 1, }) config["providers"]["openai"]["endpoints"]["embeddings"] = "https://openai.example/embeddings" config["providers"]["jina"]["endpoints"]["embeddings"] = "https://jina.example/embeddings" @@ -445,10 +447,12 @@ def test_embeddings_rejects_unsupported_configured_provider(configured_ai, monke @pytest.fixture def google_transport(monkeypatch): - calls = {"clients": [], "requests": [], "closed": []} + calls = {"clients": [], "requests": [], "closed": [], "async_closed": []} - def generate_content(**kwargs): + async def generate_content(**kwargs): calls["requests"].append(kwargs) + if "handler" in calls: + return await calls["handler"](**kwargs) if "response" in calls: return calls["response"] return SimpleNamespace(candidates=[SimpleNamespace( @@ -456,10 +460,13 @@ def generate_content(**kwargs): content=SimpleNamespace(parts=[SimpleNamespace(text='{"name":"Synthetic"}')]), )]) + async def aclose(): + calls["async_closed"].append(True) + def client(**kwargs): calls["clients"].append(kwargs) return SimpleNamespace( - models=SimpleNamespace(generate_content=generate_content), + aio=SimpleNamespace(models=SimpleNamespace(generate_content=generate_content), aclose=aclose), close=lambda: calls["closed"].append(True), ) @@ -469,6 +476,7 @@ def client(**kwargs): GenerateContentConfig=lambda **kwargs: kwargs, Tool=lambda **kwargs: kwargs, UrlContext=lambda: {}, + UrlRetrievalStatus=SimpleNamespace(URL_RETRIEVAL_STATUS_SUCCESS="URL_RETRIEVAL_STATUS_SUCCESS"), ) monkeypatch.setattr(gemini, "_get_genai", lambda: ( SimpleNamespace(Client=client), types, SimpleNamespace(ClientError=type("ClientError", (Exception,), {})), @@ -509,17 +517,27 @@ def test_google_retrieval_uses_config_at_sdk_boundary(configured_ai, google_tran assert len(google_transport["closed"]) == 1 -def test_google_packaged_defaults_omit_temperature(google_transport, monkeypatch): +@pytest.mark.parametrize("model,expected_model,thinking_level", [ + (None, "gemini-3.5-flash", "minimal"), + ("gemini-3.6-flash", "gemini-3.6-flash", "minimal"), + ("gemini-3.8-flash", "gemini-3.8-flash", "low"), +]) +def test_google_packaged_defaults_omit_temperature(google_transport, monkeypatch, model, expected_model, thinking_level): monkeypatch.delenv("WRANGLES_AI_CONFIG", raising=False) ai_config.clear_cache() try: result = gemini.GeminiURLContextClient(api_key="fake-key").retrieve( - "https://product.example/one", output_format="json", + "https://product.example/one", output_format="json", model_id=model, ) + assert google_transport["requests"][0]["model"] == expected_model assert "temperature" not in google_transport["requests"][0]["config"] + assert google_transport["requests"][0]["config"]["thinking_config"] == {"thinking_level": thinking_level} + assert google_transport["clients"][0]["http_options"]["timeout"] == 10000 + assert google_transport["clients"][0]["http_options"]["retry_options"]["attempts"] == 1 assert result["error"] is None assert result["extracted_content"] == {"name": "Synthetic"} assert google_transport["closed"] == [True] + assert google_transport["async_closed"] == [True] finally: ai_config.clear_cache() @@ -604,11 +622,14 @@ def test_google_forwards_model_tuning_defaults(configured_ai, google_transport): def test_google_closes_client_when_request_fails(configured_ai, monkeypatch): closed = [] - def generate_content(**kwargs): + async def generate_content(**kwargs): raise RuntimeError("synthetic request failure") + async def aclose(): + closed.append("async") + client = SimpleNamespace( - models=SimpleNamespace(generate_content=generate_content), + aio=SimpleNamespace(models=SimpleNamespace(generate_content=generate_content), aclose=aclose), close=lambda: closed.append(True), ) types = SimpleNamespace( @@ -622,7 +643,7 @@ def generate_content(**kwargs): result = gemini.GeminiURLContextClient(api_key="fake-key").retrieve("https://product.example/one") assert result["status"] == "Failure" assert "synthetic request failure" in result["error"] - assert closed == [True] + assert closed == ["async", True] @pytest.mark.parametrize("direct_client", [False, True]) @@ -691,6 +712,7 @@ def test_google_sdk_builds_configured_request_url(configured_ai, monkeypatch, mo import httpx from google import genai from google.genai import errors, types + from google.genai import _api_client config, save = configured_ai base_url = "https://proxy.example/google" if custom_endpoint else "https://generativelanguage.googleapis.com" @@ -699,11 +721,12 @@ def test_google_sdk_builds_configured_request_url(configured_ai, monkeypatch, mo config["operations"]["search.retrieve_link_content"]["defaults"]["api_version"] = version config["operations"]["search.retrieve_link_content"]["defaults"].pop("temperature") config["providers"]["google"]["models"]["gemini-3.8-flash"]["defaults"]["temperature"] = 0.17 + config["providers"]["google"]["models"]["gemini-3.8-flash"]["defaults"]["max_output_tokens"] = 500 save() requests_sent = [] clients = [] - def send(client, request, **kwargs): + async def send(client, request, **kwargs): requests_sent.append(request) return httpx.Response(200, request=request, json={ "candidates": [{"content": {"parts": [{"text": '{"name":"Synthetic"}'}]}}], @@ -714,23 +737,169 @@ def client(**kwargs): clients.append(result) return result - monkeypatch.setattr(httpx.Client, "send", send) + monkeypatch.setattr(httpx.AsyncClient, "send", send) + monkeypatch.setattr(_api_client, "has_aiohttp", False) monkeypatch.setattr(gemini, "_get_genai", lambda: (SimpleNamespace(Client=client), types, errors)) try: result = gemini.GeminiURLContextClient(api_key="fake-key").retrieve( "https://product.example/one", model_id=model, output_format="json", + thinking_level="low", request_timeout_seconds=2, seed=19, maxOutputTokens=321, ) assert result["error"] is None assert result["extracted_content"] == {"name": "Synthetic"} assert len(requests_sent) == 1 assert requests_sent[0].method == "POST" assert str(requests_sent[0].url) == f"{base_url}/{version}/models/gemini-3.8-flash:generateContent" - assert json.loads(requests_sent[0].content)["generationConfig"]["temperature"] == 0.17 + config_sent = json.loads(requests_sent[0].content)["generationConfig"] + assert config_sent["temperature"] == 0.17 + thinking = config_sent["thinkingConfig"] + assert thinking.get("thinkingLevel", thinking.get("thinking_level")) == "LOW" + assert config_sent["seed"] == 19 + assert config_sent["maxOutputTokens"] == 321 + assert requests_sent[0].headers["X-Server-Timeout"] == "2" finally: for client in clients: client.close() +@pytest.mark.parametrize("threads", [1, 2]) +def test_google_deadline_cancels_slow_url_without_losing_other_rows(configured_ai, google_transport, threads): + cancelled = [] + + async def generate_content(**request): + if "/slow>" in request["contents"]: + try: + await asyncio.Event().wait() + finally: + cancelled.append(True) + return SimpleNamespace(candidates=[SimpleNamespace( + url_context_metadata=SimpleNamespace(url_metadata=[SimpleNamespace( + retrieved_url="https://product.example/fast", + url_retrieval_status="URL_RETRIEVAL_STATUS_SUCCESS", + )]), + content=SimpleNamespace(parts=[SimpleNamespace(text='{"name":"Fast"}')]), + )]) + + google_transport["handler"] = generate_content + results = search.retrieve_link_content( + ["https://product.example/slow", "https://product.example/fast"], + client_config={"api_key": "fake-key"}, threads=threads, + request_timeout_seconds=0.02, + ) + assert results[0]["status"] == "Failure" + assert results[0]["error"] == "Timeout: URL retrieval exceeded 0.02 seconds." + assert results[0]["extracted_content"] is None + assert results[1]["status"] == "Success" + assert results[1]["extracted_content"] == {"name": "Fast"} + assert cancelled == [True] + assert google_transport["closed"] == google_transport["async_closed"] == [True, True] + + +def test_google_direct_client_works_inside_running_event_loop(configured_ai, google_transport): + async def retrieve(): + return gemini.GeminiURLContextClient(api_key="fake-key").retrieve( + "https://product.example/one", output_format="json", thinking_level="minimal", + ) + + assert asyncio.run(retrieve())["extracted_content"] == {"name": "Synthetic"} + assert google_transport["async_closed"] == [True] + + +@pytest.mark.parametrize("key", ["thinking_config", "thinkingConfig"]) +def test_google_raw_thinking_configuration_replaces_level(configured_ai, google_transport, key): + gemini.GeminiURLContextClient(api_key="fake-key").retrieve( + "https://product.example/one", **{key: {"thinking_budget": 0}}, + ) + config = google_transport["requests"][0]["config"] + assert config[key] == {"thinking_budget": 0} + assert "thinking_level" not in config[key] + assert ("thinkingConfig" if key == "thinking_config" else "thinking_config") not in config + + +@pytest.mark.parametrize("key", ["http_options", "httpOptions", "system_instruction", "systemInstruction", "tools", + "response_mime_type", "responseMimeType", "response_modalities", "responseModalities"]) +def test_google_kwargs_cannot_override_retrieval_controls(configured_ai, google_transport, key): + with pytest.raises(ValueError, match="Retrieval manages"): + gemini.GeminiURLContextClient(api_key="fake-key").retrieve("https://product.example/one", **{key: {}}) + assert google_transport["clients"] == [] + + +def test_google_http_timeout_has_useful_row_error(configured_ai, google_transport): + import httpx + + async def timeout(**request): + raise httpx.ReadTimeout("") + + google_transport["handler"] = timeout + result = gemini.GeminiURLContextClient(api_key="fake-key").retrieve("https://product.example/one") + assert result["status"] == "Failure" + assert result["error"].startswith("Timeout:") + assert google_transport["closed"] == google_transport["async_closed"] == [True] + + +def test_google_deadline_includes_sdk_retry_backoff(configured_ai, monkeypatch): + import httpx + from google import genai + from google.genai import _api_client, errors, types + + requests_sent = [] + clients = [] + + async def send(client, request, **kwargs): + requests_sent.append(request) + return httpx.Response(503, request=request, json={ + "error": {"code": 503, "message": "synthetic overload", "status": "UNAVAILABLE"}, + }) + + def client(**kwargs): + result = genai.Client(vertexai=False, **kwargs) + clients.append(result) + return result + + monkeypatch.setattr(_api_client, "has_aiohttp", False) + monkeypatch.setattr(httpx.AsyncClient, "send", send) + monkeypatch.setattr(gemini, "_get_genai", lambda: (SimpleNamespace(Client=client), types, errors)) + result = gemini.GeminiURLContextClient(api_key="fake-key").retrieve( + "https://product.example/one", request_timeout_seconds=0.05, + ) + assert result["error"] == "Timeout: URL retrieval exceeded 0.05 seconds." + assert result["status"] == "Failure" + assert len(requests_sent) == 1 + assert all(client._api_client._async_httpx_client.is_closed for client in clients) + assert all(client._api_client._httpx_client.is_closed for client in clients) + + +def test_google_invalid_kwargs_fail_before_creating_transport(configured_ai, google_transport, monkeypatch): + from google.genai import types + + genai, _, errors = gemini._get_genai() + monkeypatch.setattr(gemini, "_get_genai", lambda: (genai, types, errors)) + result = gemini.GeminiURLContextClient(api_key="fake-key").retrieve( + "https://product.example/one", misspelled_option=True, + ) + assert result["status"] == "Failure" + assert "Invalid Gemini generation options" in result["error"] + assert "misspelled_option" in result["error"] + assert google_transport["clients"] == [] + + +def test_google_thought_summaries_do_not_pollute_extracted_json(configured_ai, google_transport): + google_transport["response"] = SimpleNamespace(candidates=[SimpleNamespace( + url_context_metadata=None, + content=SimpleNamespace(parts=[ + SimpleNamespace(text="A thought summary", thought=True), + SimpleNamespace(text=None), + SimpleNamespace(text='{"name":"Synthetic"}', thought=False), + ]), + )]) + result = gemini.GeminiURLContextClient(api_key="fake-key").retrieve( + "https://product.example/one", output_format="json", + thinking_config={"thinking_level": "minimal", "include_thoughts": True}, + ) + assert result["extracted_content"] == {"name": "Synthetic"} + assert result["error"] is None + + def test_google_retrieval_rejects_unsupported_provider(configured_ai, monkeypatch): monkeypatch.setattr(ai_config, "resolve", lambda *args, **kwargs: {"provider": "anthropic"}) with pytest.raises(ValueError, match="only the 'google' provider"): diff --git a/tests/test_ai_config.py b/tests/test_ai_config.py index 077d0b1a..59a6ff3c 100644 --- a/tests/test_ai_config.py +++ b/tests/test_ai_config.py @@ -39,10 +39,14 @@ def test_packaged_operation_defaults_and_model_lifecycle(): assert extraction["request_timeout_seconds"] == 12 assert extraction["cache"]["ttl_seconds"] == 3600 assert ai_config.resolve("embeddings")["model"] == "text-embedding-3-small" - assert ai_config.resolve("search.retrieve_link_content")["model"] == "gemini-3.8-flash" + retrieval = ai_config.resolve("search.retrieve_link_content") + assert retrieval["model"] == "gemini-3.5-flash" + assert retrieval["thinking_level"] == "minimal" + assert retrieval["request_timeout_seconds"] == 10 + assert retrieval["retries"] == 0 for operation in config["operations"]: explicit_model = "org/task-model" if config["operations"][operation].get("requires_model") else None - assert ai_config.resolve(operation, model=explicit_model)["retries"] == 1 + assert ai_config.resolve(operation, model=explicit_model)["retries"] == (0 if operation == "search.retrieve_link_content" else 1) assert config["providers"]["anthropic"]["models"] == {} diff --git a/tests/test_search_retrieve_templates.py b/tests/test_search_retrieve_templates.py index 208b7fe9..ec1edbb1 100644 --- a/tests/test_search_retrieve_templates.py +++ b/tests/test_search_retrieve_templates.py @@ -4,6 +4,7 @@ import threading from types import SimpleNamespace +import jsonschema import pandas as pd import pytest import yaml @@ -260,6 +261,87 @@ def test_direct_python_retrieval_preserves_literal_prompt_and_return_shape(retri }] +@pytest.mark.parametrize("caller", ["python", "recipe"]) +def test_retrieval_options_reach_google_with_aligned_url_prompts(monkeypatch, caller): + calls = [] + + def retrieve(self, url, prompt, output_format, policy, **kwargs): + calls.append({"url": url, "prompt": prompt, "policy": policy, "options": kwargs}) + return _response(url, prompt) + + monkeypatch.setattr(gemini.GeminiURLContextClient, "_retrieve", retrieve) + urls = [f"https://product.example/{number}" for number in range(3)] + options = { + "thinking_level": "low", + "request_timeout_seconds": 1.5, + "max_output_tokens": 256, + "response_schema": { + "type": "object", "properties": {"title": {"type": "string"}}, + }, + } + if caller == "python": + results = search.retrieve_link_content( + urls, prompt="Keep {{ details }} literal", threads=1, + client_config={"api_key": "fake-key"}, **options, + ) + prompts = ["Keep {{ details }} literal"] * 3 + else: + result = _run(pd.DataFrame({ + "URL": [urls[:2], [], urls[2:]], + "details": ["first row", "blank row", "last row"], + }), prompt="Verify {{ details }}", threads=1, **options) + assert [len(cell) for cell in result["results"]] == [2, 0, 1] + results = [item for cell in result["results"] for item in cell] + prompts = ["Verify first row", "Verify first row", "Verify last row"] + + assert [item["retrieved_url"] for item in results] == urls + assert [item["extracted_content"]["prompt"] for item in results] == prompts + assert [call["url"] for call in calls] == urls + assert [call["prompt"] for call in calls] == prompts + for call in calls: + assert call["policy"]["thinking_level"] == "low" + assert call["policy"]["request_timeout_seconds"] == 1.5 + assert call["options"] == { + "max_output_tokens": 256, "response_schema": options["response_schema"], + } + + +@pytest.mark.parametrize("caller", ["python", "recipe"]) +@pytest.mark.parametrize("options", [ + {}, + {"thinking_level": None, "request_timeout_seconds": None}, + {"thinking_level": "minimal"}, + {"request_timeout_seconds": 2.5}, + {"thinking_level": "low", "request_timeout_seconds": 4, "max_output_tokens": 128}, +]) +def test_custom_retrievers_receive_only_explicit_options(retrieval_calls, caller, options): + url = "https://product.example/one" + if caller == "python": + result = search.retrieve_link_content( + url, prompt="Page title", model_id="custom-model", **options, + ) + else: + result = _run(pd.DataFrame({"URL": [url]}), + prompt="Page title", model_id="custom-model", **options)["results"].iloc[0][0] + assert result["extracted_content"] == {"prompt": "Page title"} + assert retrieval_calls["requests"] == [{ + "url": url, "prompt": "Page title", "model_id": "custom-model", "output_format": "json", + **{key: value for key, value in options.items() if value is not None}, + }] + + +def test_retrieval_schema_accepts_runtime_overrides_and_additional_sdk_options(): + schema = yaml.safe_load(recipe_search.retrieve_link_content.__doc__) + jsonschema.Draft7Validator.check_schema(schema) + jsonschema.validate({ + "input": "URL", "output": "results", "prompt": "Verify {{ details }}", + "thinking_level": "minimal", "request_timeout_seconds": 10, + "max_output_tokens": 256, + "response_schema": {"type": "object", "properties": {"title": {"type": "string"}}}, + }, schema) + assert {"thinking_level", "request_timeout_seconds"} <= schema["properties"].keys() + + @pytest.mark.parametrize("threads", [None, 2]) def test_row_prompts_share_one_bounded_parallel_batch(monkeypatch, threads): calls = [] @@ -334,13 +416,13 @@ def retrieve(**kwargs): def test_row_prompts_and_options_reach_real_google_sdk_without_network(monkeypatch, output_format, mime_type): import httpx from google import genai - from google.genai import errors, types + from google.genai import _api_client, errors, types requests = [] client_options = [] clients = [] - def send(client, request, **kwargs): + async def send(client, request, **kwargs): requests.append(request) return httpx.Response(200, request=request, json={ "candidates": [{"content": {"parts": [{"text": '{"name":"synthetic"}'}]}}], @@ -352,14 +434,16 @@ def client(**kwargs): clients.append(instance) return instance - monkeypatch.setattr(httpx.Client, "send", send) + monkeypatch.setattr(_api_client, "has_aiohttp", False) + monkeypatch.setattr(httpx.AsyncClient, "send", send) monkeypatch.setattr(gemini, "_get_genai", lambda: (SimpleNamespace(Client=client), types, errors)) try: result = _run(pd.DataFrame({ "URL": ["https://product.example/one", "https://product.example/two"], "details": ["brass fitting", "steel bearing"], }), prompt="Verify {{ details }}", output_format=output_format, threads=1, - model_id="models/private-google-model") + model_id="models/private-google-model", thinking_level="low", + request_timeout_seconds=1.5, max_output_tokens=256, seed=17) bodies = [json.loads(request.content) for request in requests] assert len(bodies) == 2 for request, body, details in zip(requests, bodies, ["brass fitting", "steel bearing"]): @@ -367,10 +451,14 @@ def client(**kwargs): assert body["systemInstruction"]["parts"][0]["text"].startswith(f"Verify {details}\n\n") assert body["generationConfig"]["responseMimeType"] == mime_type assert body["generationConfig"]["temperature"] == 0.17 + thinking = body["generationConfig"]["thinkingConfig"] + assert thinking.get("thinkingLevel", thinking.get("thinking_level")) == "LOW" + assert body["generationConfig"]["maxOutputTokens"] == 256 + assert body["generationConfig"]["seed"] == 17 assert body["tools"] == [{"urlContext": {}}] assert "https://product.example/one" in bodies[0]["contents"][0]["parts"][0]["text"] assert "https://product.example/two" in bodies[1]["contents"][0]["parts"][0]["text"] - assert all(options["http_options"].timeout == 3250 for options in client_options) + assert all(options["http_options"].timeout == 1500 for options in client_options) assert all(options["http_options"].retry_options.attempts == 1 for options in client_options) assert all(cell[0]["error"] is None for cell in result["results"]) finally: diff --git a/wrangles/ai_defaults.yml b/wrangles/ai_defaults.yml index 8b9a40a0..236004f2 100644 --- a/wrangles/ai_defaults.yml +++ b/wrangles/ai_defaults.yml @@ -82,11 +82,38 @@ providers: documentation: model_cards: https://deepmind.google/models/model-cards/ models: - gemini-3.8-flash: + gemini-3.5-flash: status: active applications: [url_retrieval] default_for: [search.retrieve_link_content] - defaults: {} + defaults: + thinking_level: minimal + supported_values: + thinking_level: [minimal, low, medium, high] + gemini-3.6-flash: + status: active + applications: [url_retrieval] + default_for: [] + defaults: + thinking_level: minimal + supported_values: + thinking_level: [minimal, low, medium, high] + gemini-3-flash-preview: + status: active + applications: [url_retrieval] + default_for: [] + defaults: + thinking_level: minimal + supported_values: + thinking_level: [minimal, low, medium, high] + gemini-3.8-flash: + status: active + applications: [url_retrieval] + default_for: [] + defaults: + thinking_level: low + supported_values: + thinking_level: [low, medium, high] jina: endpoints: embeddings: https://api.jina.ai/v1/embeddings @@ -235,8 +262,8 @@ operations: defaults: api_version: v1beta default_concurrency: 10 - request_timeout_seconds: 45 - retries: 1 + request_timeout_seconds: 10 + retries: 0 huggingface: provider: huggingface protocol: hf_inference diff --git a/wrangles/clients/gemini.py b/wrangles/clients/gemini.py index eb6c78d9..66965a76 100644 --- a/wrangles/clients/gemini.py +++ b/wrangles/clients/gemini.py @@ -1,3 +1,6 @@ +import asyncio +import concurrent.futures +import math import os from typing import Optional, Dict, Any from ..utils import LazyLoader as _LazyLoader @@ -7,6 +10,7 @@ genai = _LazyLoader("google.genai") types = _LazyLoader("google.genai.types") errors = _LazyLoader("google.genai.errors") +_httpx = _LazyLoader("httpx") def _get_genai(): """ @@ -48,7 +52,9 @@ def __init__(self, api_key: Optional[str] = None): "Missing API Key: Provide `api_key` in the recipe config or set the GOOGLE_API_KEY environment variable." ) - def retrieve(self, url: str, prompt: Optional[str] = None, model_id: str = None, output_format: str = "markdown") -> Dict[str, Any]: + def retrieve(self, url: str, prompt: Optional[str] = None, model_id: str = None, + output_format: str = "markdown", thinking_level: str = None, + request_timeout_seconds: float = None, **kwargs) -> Dict[str, Any]: """ Retrieves context from a web URL using the Gemini API. Includes thread-safe initialization, strict timeouts, and optional JSON parsing. @@ -56,10 +62,34 @@ def retrieve(self, url: str, prompt: Optional[str] = None, model_id: str = None, policy = None if url and str(url).strip(): policy = _ai_config.resolve("search.retrieve_link_content", model=model_id) + if thinking_level is not None: + policy["thinking_level"] = thinking_level + if request_timeout_seconds is not None: + policy["request_timeout_seconds"] = request_timeout_seconds _ai_config.warn_if_deprecated(policy["model"], policy["provider"]) - return self._retrieve(url, prompt, output_format, policy) + return self._retrieve(url, prompt, output_format, policy, **kwargs) - def _retrieve(self, url, prompt, output_format, policy): + @staticmethod + def _generate_with_deadline(client, timeout, **request): + """Cancel the whole SDK request, including retries, before closing it.""" + async def generate(): + try: + return await asyncio.wait_for( + client.aio.models.generate_content(**request), timeout=timeout, + ) + finally: + await client.aio.aclose() + + try: + asyncio.get_running_loop() + except RuntimeError: + return asyncio.run(generate()) + # The public client remains usable from synchronous code inside an + # existing event loop, without attempting to nest asyncio.run(). + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + return executor.submit(lambda: asyncio.run(generate())).result() + + def _retrieve(self, url, prompt, output_format, policy, **kwargs): """Retrieve one URL using the operation's already resolved policy.""" result = { "retrieved_url": url, @@ -81,6 +111,20 @@ def _retrieve(self, url, prompt, output_format, policy): if policy["protocol"] != "generate_content": raise ValueError("Google URL content retrieval requires the 'generate_content' protocol.") model_id = policy["model"] + timeout = policy["request_timeout_seconds"] + if type(timeout) not in (int, float) or not math.isfinite(timeout) or timeout <= 0: + raise ValueError("request_timeout_seconds must be a positive finite number.") + thinking_level = policy.get("thinking_level") + if thinking_level is not None and thinking_level not in ("minimal", "low", "medium", "high"): + raise ValueError("thinking_level must be minimal, low, medium, or high.") + reserved = { + "system_instruction", "systemInstruction", "tools", + "response_mime_type", "responseMimeType", "response_modalities", "responseModalities", + "http_options", "httpOptions", + } + if reserved.intersection(kwargs): + names = ", ".join(sorted(reserved.intersection(kwargs))) + raise ValueError(f"Retrieval manages {names}; use prompt, output_format, or request_timeout_seconds instead.") user_content = f"Please retrieve content from this explicitly bounded URL: <{url}>" @@ -94,6 +138,31 @@ def _retrieve(self, url, prompt, output_format, policy): system_instruction = base_prompt + "\n\nCRITICAL FORMAT RULE:\n- Strictly use Markdown.\n- Use the section names exactly as listed, with empty lines between each section." mime_type = "text/plain" + try: + generation_config = { + "system_instruction": system_instruction, + "tools": [self.types.Tool(url_context=self.types.UrlContext())], + "response_modalities": ["TEXT"], + "response_mime_type": mime_type, + **{key: policy[key] for key in ( + "temperature", "top_p", "top_k", "max_output_tokens", "stop_sequences", + ) if key in policy}, + } + if thinking_level is not None: + generation_config["thinking_config"] = {"thinking_level": thinking_level} + # SDK camelCase aliases override snake_case defaults too. Replace + # thinking_config as a whole to avoid combining a level and budget. + for key in tuple(generation_config): + first, *rest = key.split("_") + alias = first + "".join(word.capitalize() for word in rest) + if alias != key and alias in kwargs: + generation_config.pop(key) + generation_config.update(kwargs) + generation_config = self.types.GenerateContentConfig(**generation_config) + except Exception as e: + result["error"] = f"Invalid Gemini generation options: {e}" + return result + try: # The SDK expects milliseconds; configuration stores seconds. client = self.genai.Client( @@ -101,7 +170,7 @@ def _retrieve(self, url, prompt, output_format, policy): http_options=self.types.HttpOptions( base_url=policy.get("endpoints", {}).get("base_url"), api_version=policy.get("api_version", "v1beta"), - timeout=policy["request_timeout_seconds"] * 1000, + timeout=timeout * 1000, retry_options=self.types.HttpRetryOptions( attempts=policy["retries"] + 1, http_status_codes=[408, 429, 500, 502, 503, 504], @@ -113,18 +182,11 @@ def _retrieve(self, url, prompt, output_format, policy): return result try: - response = client.models.generate_content( + response = self._generate_with_deadline( + client, timeout, model=model_id, contents=user_content, - config=self.types.GenerateContentConfig( - system_instruction=system_instruction, - tools=[self.types.Tool(url_context=self.types.UrlContext())], - response_modalities=["TEXT"], - response_mime_type=mime_type, - **{key: policy[key] for key in ( - "temperature", "top_p", "top_k", "max_output_tokens", "stop_sequences", - ) if key in policy}, - ) + config=generation_config, ) cand = response.candidates[0] if response.candidates else None @@ -154,7 +216,10 @@ def _retrieve(self, url, prompt, output_format, policy): result["error"] += f". {cand.finish_message}" result["status"] = "Failure" else: - full_text = "\n".join([part.text for part in response.candidates[0].content.parts]) + full_text = "\n".join( + part.text for part in cand.content.parts + if part.text and not getattr(part, "thought", False) + ) if output_format.lower() == "json": import json @@ -174,6 +239,9 @@ def _retrieve(self, url, prompt, output_format, policy): else: result["extracted_content"] = full_text + except (asyncio.TimeoutError, TimeoutError, _httpx.TimeoutException): + result["status"] = "Failure" + result["error"] = f"Timeout: URL retrieval exceeded {timeout:g} seconds." except self.errors.ClientError as e: result["status"] = "Failure" error_str = str(e) @@ -193,6 +261,10 @@ def _retrieve(self, url, prompt, output_format, policy): else: result["error"] = f"Unexpected Error: {error_str}" finally: - client.close() + try: + client.close() + except Exception as e: + result["status"] = "Failure" + result["error"] = result["error"] or f"Failed to close retrieval client: {e}" return result diff --git a/wrangles/recipe_wrangles/search.py b/wrangles/recipe_wrangles/search.py index 51d761ae..c7d8f1bd 100644 --- a/wrangles/recipe_wrangles/search.py +++ b/wrangles/recipe_wrangles/search.py @@ -314,12 +314,15 @@ def retrieve_link_content( prompt: str | None = None, model_id: str = None, output_format: str = "json", - threads: int = None + threads: int = None, + thinking_level: str | None = None, + request_timeout_seconds: float | None = None, + **kwargs, ) -> _pd.DataFrame: """ type: object - description: Retrieves targeted content from web pages using LLM URL extraction. Can optionally output a second column containing a clean, human-readable text summary of the retrieved data. - additionalProperties: false + description: Retrieves targeted content from web pages using LLM URL extraction. Can optionally output a second column containing a clean, human-readable text summary of the retrieved data. Additional options are passed to Gemini GenerateContentConfig. + additionalProperties: true required: - input - output @@ -366,6 +369,18 @@ def retrieve_link_content( threads: type: integer description: Number of concurrent threads for parallel processing. Defaults to the AI configuration. + thinking_level: + type: string + enum: + - minimal + - low + - medium + - high + description: Gemini thinking level. Defaults to the configured value, packaged as minimal. Supported levels depend on the selected model. + request_timeout_seconds: + type: number + exclusiveMinimum: 0 + description: Total deadline in seconds for each URL, including retries. Defaults to the configured value, packaged as 10 seconds. A batch may take longer when URLs run in several waves. """ if output is None: output = input @@ -432,7 +447,10 @@ def _to_url_list(v) -> list[str]: client_config=client_config, model_id=model_id, output_format=output_format, - threads=threads + threads=threads, + thinking_level=thinking_level, + request_timeout_seconds=request_timeout_seconds, + **kwargs, ) out_cells_dict, out_cells_text = [], [] diff --git a/wrangles/search.py b/wrangles/search.py index a7e47a81..cc88926b 100644 --- a/wrangles/search.py +++ b/wrangles/search.py @@ -1,4 +1,5 @@ import concurrent.futures as _futures +import math as _math # Import our client factory from .clients import get_client as _get_client @@ -72,18 +73,24 @@ def retrieve_link_content( prompt: str | None = None, model_id: str = None, output_format: str = "json", - threads: int = None + threads: int = None, + thinking_level: str | None = None, + request_timeout_seconds: float | None = None, + **kwargs, ) -> dict | list: """ Retrieve formatted content from web URLs using a specified client. - Omitted model and concurrency settings are resolved from the AI configuration. + Omitted model, thinking, concurrency, and deadline settings use AI configuration. The prompt is shared literally across URLs; column templates are recipe-only. + Additional keyword arguments are Gemini GenerateContentConfig options. """ is_scalar = not isinstance(urls, list) urls = [urls] if is_scalar else urls results = _retrieve_link_content( [(url, prompt) for url in urls], client=client, client_config=client_config, model_id=model_id, output_format=output_format, threads=threads, + thinking_level=thinking_level, request_timeout_seconds=request_timeout_seconds, + **kwargs, ) return results[0] if is_scalar else results @@ -95,9 +102,28 @@ def _retrieve_link_content( model_id: str = None, output_format: str = "json", threads: int = None, + thinking_level: str | None = None, + request_timeout_seconds: float | None = None, + **kwargs, ) -> list: """Retrieve ordered URL/prompt pairs with one policy and bounded worker pool.""" + if thinking_level is not None and thinking_level not in ("minimal", "low", "medium", "high"): + raise ValueError("thinking_level must be minimal, low, medium, or high.") + if request_timeout_seconds is not None and ( + isinstance(request_timeout_seconds, bool) + or not isinstance(request_timeout_seconds, (int, float)) + or not _math.isfinite(request_timeout_seconds) + or request_timeout_seconds <= 0 + ): + raise ValueError("request_timeout_seconds must be a positive finite number.") policy = _ai_config.resolve("search.retrieve_link_content", model=model_id) + overrides = { + key: value for key, value in ( + ("thinking_level", thinking_level), + ("request_timeout_seconds", request_timeout_seconds), + ) if value is not None + } + policy.update(overrides) if policy["provider"] != "google": raise ValueError("URL content retrieval currently supports only the 'google' provider.") if policy["protocol"] != "generate_content": @@ -112,8 +138,11 @@ def _retrieve_link_content( def retrieve(request): url, prompt = request if isinstance(retriever, _GeminiURLContextClient): - return retriever._retrieve(url, prompt, output_format, policy) - return retriever.retrieve(url=url, prompt=prompt, model_id=model_id, output_format=output_format) + return retriever._retrieve(url, prompt, output_format, policy, **kwargs) + return retriever.retrieve( + url=url, prompt=prompt, model_id=model_id, output_format=output_format, + **overrides, **kwargs, + ) with _futures.ThreadPoolExecutor(max_workers=threads) as executor: return list(executor.map(retrieve, requests))