From d6d5afa382a49d35b264e3c2ab46bef5d601ddcf Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Mon, 28 Sep 2026 16:42:44 +0200 Subject: [PATCH 1/4] Python: add Gemini Embedding 2 and per-operation vector options --- docs/specs/feature-usage-bit-registry.md | 2 +- python/.env.example | 4 +- python/packages/core/AGENTS.md | 8 + .../packages/core/agent_framework/_vectors.py | 193 ++++- .../core/agent_framework/gemini/__init__.py | 4 +- .../core/agent_framework/gemini/__init__.pyi | 8 +- .../core/tests/core/test_gemini_namespace.py | 6 +- .../packages/core/tests/core/test_vectors.py | 201 +++++ python/packages/gemini/AGENTS.md | 26 +- python/packages/gemini/README.md | 90 ++- .../gemini/agent_framework_gemini/__init__.py | 10 +- .../agent_framework_gemini/_chat_client.py | 151 +--- .../_embedding_client.py | 408 ++++++++++ .../agent_framework_gemini/_sdk_client.py | 127 ++++ python/packages/gemini/pyproject.toml | 4 +- .../gemini/test_gemini_embedding_client.py | 696 ++++++++++++++++++ .../gemini/tests/test_gemini_client.py | 54 +- .../02-agents/providers/gemini/README.md | 19 +- .../providers/gemini/gemini_advanced.py | 3 +- .../providers/gemini/gemini_basic.py | 3 +- .../providers/gemini/gemini_embeddings.py | 50 ++ .../gemini/gemini_image_search_agent.py | 108 +++ .../providers/gemini/gemini_search_agent.py | 97 +++ .../gemini/gemini_with_code_execution.py | 3 +- .../gemini/gemini_with_google_maps.py | 3 +- .../gemini/gemini_with_google_search.py | 3 +- .../samples/02-agents/vector_stores/README.md | 14 + python/samples/README.md | 7 + python/uv.lock | 10 +- 29 files changed, 2112 insertions(+), 200 deletions(-) create mode 100644 python/packages/gemini/agent_framework_gemini/_embedding_client.py create mode 100644 python/packages/gemini/agent_framework_gemini/_sdk_client.py create mode 100644 python/packages/gemini/tests/gemini/test_gemini_embedding_client.py create mode 100644 python/samples/02-agents/providers/gemini/gemini_embeddings.py create mode 100644 python/samples/02-agents/providers/gemini/gemini_image_search_agent.py create mode 100644 python/samples/02-agents/providers/gemini/gemini_search_agent.py diff --git a/docs/specs/feature-usage-bit-registry.md b/docs/specs/feature-usage-bit-registry.md index 4f7ee4cd937..939ed2748ed 100644 --- a/docs/specs/feature-usage-bit-registry.md +++ b/docs/specs/feature-usage-bit-registry.md @@ -158,7 +158,7 @@ only to approved first-party endpoints. | 56 | `openai` | OpenAI clients | `agent_framework_openai` | | 57 | `anthropic` | Anthropic clients | `agent_framework_anthropic` | | 58 | `bedrock` | AWS Bedrock clients | `agent_framework_bedrock` | -| 59 | `gemini` | Gemini chat client | `agent_framework_gemini` | +| 59 | `gemini` | Gemini chat and embedding clients | `agent_framework_gemini` | | 60 | `mistral` | Mistral embedding client | `agent_framework_mistral` | | 61 | `ollama` | Ollama clients | `agent_framework_ollama` | | 62 | `claude` | Claude Agent SDK agent | `agent_framework_claude` | diff --git a/python/.env.example b/python/.env.example index b37b1787935..f9f3a953c4e 100644 --- a/python/.env.example +++ b/python/.env.example @@ -41,8 +41,8 @@ COPILOTSTUDIOAGENT__AGENTAPPID="" ANTHROPIC_API_KEY="" ANTHROPIC_MODEL="" # Google Gemini -GEMINI_API_KEY="" -GEMINI_MODEL="" +GOOGLE_API_KEY="" +GOOGLE_MODEL="" # Ollama OLLAMA_ENDPOINT="" OLLAMA_MODEL="" diff --git a/python/packages/core/AGENTS.md b/python/packages/core/AGENTS.md index af80d5bdb0a..f5023630ec7 100644 --- a/python/packages/core/AGENTS.md +++ b/python/packages/core/AGENTS.md @@ -125,6 +125,14 @@ The vector store API is experimental under the shared `VECTOR_STORES` feature ID deserializes results without interpreting thresholds or re-filtering returned scores. Connectors own scoring, filter execution, score thresholds (including provider-defined/default metrics), and paging. Use native backend execution where available, otherwise an explicit connector-local fallback or reject unsupported options +- **Embedding request options** - `upsert` accepts either flat `embeddings_options` for + all generated vector fields or `embeddings_options_by_field` keyed by logical field name, + never both. `search` and `create_vector_search_tool` accept flat `embeddings_options` + for local query generation. Core supplies declared field dimensions and rejects + conflicting values before embedding; search ignores these options when a + precomputed vector is supplied. `create_upsert_tool` and + `VectorCollectionContextProvider` forward the same operation-specific settings + to their generated tools. - **`create_vector_search_tool`** - Creates an agent tool from any `SupportsVectorSearch` implementation - **`create_upsert_tool` / `create_get_tool` / `create_delete_tool`** - Create agent tools for collection CRUD; upsert and delete require approval by default, while get does not. Auto-generated keys are omitted from upsert diff --git a/python/packages/core/agent_framework/_vectors.py b/python/packages/core/agent_framework/_vectors.py index 4849a38da0a..54273b5d3b2 100644 --- a/python/packages/core/agent_framework/_vectors.py +++ b/python/packages/core/agent_framework/_vectors.py @@ -46,7 +46,7 @@ from ._sessions import AgentSession, ContextProvider, HistoryProvider, SessionContext from ._telemetry import FeatureIndex, mark_feature_used from ._tools import ApprovalMode, FunctionTool -from ._types import Content, EmbeddingGenerationOptions, Message +from ._types import Content, Message from ._vector_filters import ( Filter, FilterExpression, @@ -358,6 +358,55 @@ def __init__( object.__setattr__(self, "provider_annotations", _copy_provider_annotations(provider_annotations)) +def _prepare_embedding_options( + field: VectorStoreField, + embeddings_options: Mapping[str, Any] | None, +) -> dict[str, Any]: + options = deepcopy(dict(embeddings_options)) if embeddings_options is not None else {} + if field.dimensions is not None: + requested = options.get("dimensions") + if "dimensions" in options and ( + not isinstance(requested, int) or isinstance(requested, bool) or requested != field.dimensions + ): + raise ValueError( + f"embeddings_options['dimensions'] for vector field '{field.name}' " + f"must match its declared dimensions ({field.dimensions})." + ) + options["dimensions"] = field.dimensions + return options + + +def _prepare_upsert_embedding_options( + vector_fields: Sequence[VectorStoreField], + *, + record_count: int, + embeddings_options: Mapping[str, Any] | None, + embeddings_options_by_field: Mapping[str, Mapping[str, Any]] | None, +) -> dict[str, dict[str, Any]]: + if embeddings_options is not None and embeddings_options_by_field is not None: + raise ValueError("Provide either embeddings_options or embeddings_options_by_field, not both.") + if not vector_fields or not record_count: + if embeddings_options is not None or embeddings_options_by_field is not None: + raise ValueError("Embedding options require records and generated vector fields.") + return {} + if embeddings_options_by_field is not None: + if not isinstance(embeddings_options_by_field, Mapping): + raise TypeError("embeddings_options_by_field must be a mapping.") + if any(not isinstance(name, str) for name in embeddings_options_by_field): + raise TypeError("embeddings_options_by_field keys must be vector field names.") + unknown = sorted(set(embeddings_options_by_field) - {field.name for field in vector_fields}) + if unknown: + raise ValueError(f"Embedding options provided for non-generated vector field(s): {', '.join(unknown)}.") + for name, options in embeddings_options_by_field.items(): + if not isinstance(options, Mapping): + raise TypeError(f"embeddings_options_by_field[{name!r}] must be a mapping.") + return { + field.name: _prepare_embedding_options(field, embeddings_options_by_field.get(field.name)) + for field in vector_fields + } + return {field.name: _prepare_embedding_options(field, embeddings_options) for field in vector_fields} + + @experimental(feature_id=ExperimentalFeature.VECTOR_STORES) @dataclass(frozen=True, slots=True, init=False) class VectorStoreCollectionDefinition: @@ -916,6 +965,8 @@ async def serialize( records: ModelT | Sequence[ModelT], *, generate_vectors: GenerateVectors = True, + embeddings_options: Mapping[str, Any] | None = None, + embeddings_options_by_field: Mapping[str, Mapping[str, Any]] | None = None, context: Mapping[str, Any] | None = None, ) -> Any: """Serialize one or more application records for the backing store. @@ -930,6 +981,12 @@ async def serialize( records: One application record or a sequence of records. generate_vectors: Whether to generate all vector fields, preserve all supplied values, or generate only the vector fields named in a sequence. Generated values overwrite supplied values. + embeddings_options: Provider embedding options applied to every generated vector field. + Dimensions come from the vector field and conflicting values raise an error. + Mutually exclusive with ``embeddings_options_by_field``. + embeddings_options_by_field: Provider options keyed by logical vector field name. + Only fields selected by ``generate_vectors`` may appear; omitted fields receive + their declared dimensions only. context: Connector-specific serialization context. Raises: @@ -943,8 +1000,16 @@ async def serialize( dict_records = [self._serialize_record_to_dict(record) for record in input_records] vector_fields = self._resolve_vector_fields_to_generate(generate_vectors) + field_options = _prepare_upsert_embedding_options( + vector_fields, + record_count=len(dict_records), + embeddings_options=embeddings_options, + embeddings_options_by_field=embeddings_options_by_field, + ) if vector_fields: - await self._add_vectors_to_records(dict_records, vector_fields=vector_fields) + await self._add_vectors_to_records( + dict_records, vector_fields=vector_fields, embeddings_options_by_field=field_options + ) dimension_fields = tuple((field.storage_name or field.name, field) for field in self.definition.vector_fields) for record_index, record in enumerate(dict_records): for storage_name, field in dimension_fields: @@ -1025,6 +1090,7 @@ async def _add_vectors_to_records( records: Sequence[dict[str, Any]], *, vector_fields: Sequence[VectorStoreField], + embeddings_options_by_field: Mapping[str, Mapping[str, Any]], ) -> None: if not records: return @@ -1045,10 +1111,9 @@ async def _add_vectors_to_records( raise ValueError( f"Vector field '{field.name}' cannot be embedded because at least one value is missing." ) - options: EmbeddingGenerationOptions = {} - if field.dimensions is not None: - options["dimensions"] = field.dimensions - embeddings = await embedding_generator.get_embeddings(values, options=options) + embeddings = await embedding_generator.get_embeddings( + values, options=embeddings_options_by_field[field.name] + ) if len(embeddings) != len(records): raise IntegrationInvalidResponseException( f"Embedding client returned {len(embeddings)} vectors for {len(records)} records." @@ -1276,6 +1341,8 @@ async def upsert( records: Sequence[ModelT], *, generate_vectors: GenerateVectors = True, + embeddings_options: Mapping[str, Any] | None = None, + embeddings_options_by_field: Mapping[str, Mapping[str, Any]] | None = None, operation_options: Mapping[str, Any] | None = None, ) -> Sequence[KeyT]: """Upsert a batch of records. @@ -1295,6 +1362,12 @@ async def upsert( records: A sequence of models. generate_vectors: Whether to generate all vector fields, preserve all supplied values, or generate only the vector fields named in a sequence. Generated values overwrite supplied values. + embeddings_options: Provider options applied to every generated vector field. + The field's dimensions are added automatically; a different supplied + ``dimensions`` value raises an error. Mutually exclusive with + ``embeddings_options_by_field``. + embeddings_options_by_field: Provider options keyed by logical vector field + name, restricted to fields selected for generation. operation_options: Store-specific operation options. Returns: @@ -1310,7 +1383,12 @@ async def upsert( if not _is_non_string_sequence(records): raise TypeError("records must be a sequence.") try: - serialized = await self.serialize(records, generate_vectors=generate_vectors) + serialized = await self.serialize( + records, + generate_vectors=generate_vectors, + embeddings_options=embeddings_options, + embeddings_options_by_field=embeddings_options_by_field, + ) store_records = list(serialized) if _is_non_string_sequence(serialized) else [serialized] keys = list(await self._inner_upsert(store_records, operation_options=operation_options)) except (TypeError, ValueError, NotImplementedError): @@ -1549,6 +1627,7 @@ async def search( *, search_type: SearchType = "vector", vector: Vector | None = None, + embeddings_options: Mapping[str, Any] | None = None, filter: FilterExpression | None = None, top: int = 3, skip: int = 0, @@ -1568,6 +1647,9 @@ async def search( values: The value to search for or vectorize. search_type: Whether to perform vector or keyword-hybrid search. vector: An optional precomputed query vector. + embeddings_options: Provider options used only for local query embedding. + Ignored when ``vector`` is provided. The selected vector field supplies + dimensions; conflicts raise an error. filter: A portable data-only filter. top: The maximum number of results. skip: The number of results to skip. @@ -1594,6 +1676,7 @@ async def search( *, search_type: Literal["vector"] = "vector", vector: Vector, + embeddings_options: Mapping[str, Any] | None = None, filter: FilterExpression | None = None, top: int = 3, skip: int = 0, @@ -1612,6 +1695,7 @@ async def search( Args: search_type: The vector search type. vector: The precomputed query vector. + embeddings_options: Ignored when supplying a precomputed vector. filter: A portable data-only filter. top: The maximum number of results. skip: The number of results to skip. @@ -1638,6 +1722,7 @@ async def search( *, search_type: SearchType = "vector", vector: Vector | None = None, + embeddings_options: Mapping[str, Any] | None = None, filter: FilterExpression | None = None, top: int = 3, skip: int = 0, @@ -1664,6 +1749,9 @@ async def search( values: The value to search for or vectorize. search_type: Whether to perform vector or keyword-hybrid search. vector: A precomputed query vector. + embeddings_options: Provider options used only for local query embedding. + Ignored when ``vector`` is provided. The selected vector field supplies + dimensions; conflicts raise an error. filter: A portable data-only filter. top: The maximum number of results. skip: The number of results to skip. @@ -1702,6 +1790,7 @@ async def search( resolved_vector = await self._generate_vector_from_values( values, vector_property_name=vector_property_name, + embeddings_options=embeddings_options, ) if resolved_vector is not None: vector_field = self.definition.try_get_vector_field(vector_property_name) @@ -1743,19 +1832,22 @@ async def _generate_vector_from_values( values: Any, *, vector_property_name: str | None, + embeddings_options: Mapping[str, Any] | None, ) -> Vector | None: vector_field = self.definition.try_get_vector_field(vector_property_name) if vector_field is None: if vector_property_name is not None: raise ValueError(f"Vector field '{vector_property_name}' was not found in the collection definition.") + if embeddings_options is not None: + raise ValueError("embeddings_options requires a vector field for local query embedding.") return None embedding_generator = vector_field.embedding_generator or self.embedding_generator if embedding_generator is None: + if embeddings_options is not None: + raise ValueError("embeddings_options requires an embedding generator for local query embedding.") return None - embedding_options: EmbeddingGenerationOptions = {} - if vector_field.dimensions is not None: - embedding_options["dimensions"] = vector_field.dimensions - embeddings = await embedding_generator.get_embeddings([values], options=embedding_options) + options = _prepare_embedding_options(vector_field, embeddings_options) + embeddings = await embedding_generator.get_embeddings([values], options=options) if len(embeddings) != 1: raise IntegrationInvalidResponseException( f"Embedding client returned {len(embeddings)} vectors for one search value." @@ -1815,6 +1907,8 @@ async def upsert( records: Sequence[ModelT], *, generate_vectors: GenerateVectors = True, + embeddings_options: Mapping[str, Any] | None = None, + embeddings_options_by_field: Mapping[str, Mapping[str, Any]] | None = None, operation_options: Mapping[str, Any] | None = None, ) -> Sequence[KeyT]: """Upsert a batch, which may partially succeed, generating embeddings by default.""" @@ -1861,6 +1955,7 @@ async def search( *, search_type: SearchType = "vector", vector: Vector | None = None, + embeddings_options: Mapping[str, Any] | None = None, filter: FilterExpression | None = None, top: int = 3, skip: int = 0, @@ -1876,6 +1971,8 @@ async def search( values: The value to search for or vectorize. search_type: Whether to perform vector or keyword-hybrid search. vector: An optional precomputed query vector. + embeddings_options: Provider options for locally embedding the query; + ignored when a precomputed vector is provided. filter: A portable data-only filter. top: The maximum number of results. skip: The number of results to skip. @@ -1902,6 +1999,7 @@ async def search( *, search_type: Literal["vector"] = "vector", vector: Vector, + embeddings_options: Mapping[str, Any] | None = None, filter: FilterExpression | None = None, top: int = 3, skip: int = 0, @@ -1916,6 +2014,7 @@ async def search( Args: search_type: The vector search type. vector: The precomputed query vector. + embeddings_options: Ignored when supplying a precomputed vector. filter: A portable data-only filter. top: The maximum number of results. skip: The number of results to skip. @@ -2200,6 +2299,8 @@ def create_upsert_tool( description: str = _DEFAULT_UPSERT_TOOL_DESCRIPTION, approval_mode: Literal["always_require", "never_require"] = "always_require", generate_vectors: GenerateVectors = True, + embeddings_options: Mapping[str, Any] | None = None, + embeddings_options_by_field: Mapping[str, Mapping[str, Any]] | None = None, filter: FilterExpression | None = None, max_batch_size: int = _DEFAULT_VECTOR_TOOL_MAX_BATCH_SIZE, ) -> FunctionTool: @@ -2217,6 +2318,10 @@ def create_upsert_tool( description: The tool description shown to the model. approval_mode: Whether the tool requires approval before invocation. generate_vectors: Which vector fields the collection generates during upsert. + embeddings_options: Provider options applied to every generated vector field. + Conflicting dimensions raise an error. Mutually exclusive with + ``embeddings_options_by_field``. + embeddings_options_by_field: Provider options keyed by logical vector field name. filter: Optional fixed scope filter that every candidate record must satisfy. Unsupported filter operators fail closed before embedding or writing. max_batch_size: Maximum records accepted in one invocation. @@ -2225,6 +2330,14 @@ def create_upsert_tool( A function tool accepting a non-empty ``records`` array. """ _validate_vector_tool_max_batch_size(max_batch_size) + if embeddings_options is not None and embeddings_options_by_field is not None: + raise ValueError("Provide either embeddings_options or embeddings_options_by_field, not both.") + configured_embeddings_options = deepcopy(dict(embeddings_options)) if embeddings_options is not None else None + configured_by_field: dict[str, dict[str, Any]] | None = None + if embeddings_options_by_field is not None: + if not isinstance(embeddings_options_by_field, Mapping): + raise TypeError("embeddings_options_by_field must be a mapping.") + configured_by_field = {name: deepcopy(dict(options)) for name, options in embeddings_options_by_field.items()} configured_filter = _prepare_vector_tool_filter(collection, filter) async def upsert_tool(records: Any) -> dict[str, Any]: @@ -2248,7 +2361,15 @@ async def upsert_tool(records: Any) -> dict[str, Any]: if invalid_indexes: indexes = ", ".join(str(index) for index in invalid_indexes) raise ValueError(f"records at indexes {indexes} do not satisfy the configured scope filter.") - keys = await collection.upsert(decoded, generate_vectors=generate_vectors) + if configured_embeddings_options is None and configured_by_field is None: + keys = await collection.upsert(decoded, generate_vectors=generate_vectors) + else: + keys = await collection.upsert( + decoded, + generate_vectors=generate_vectors, + embeddings_options=deepcopy(configured_embeddings_options), + embeddings_options_by_field=deepcopy(configured_by_field), + ) return {"keys": [collection.key_to_json(key) for key in keys]} return FunctionTool( @@ -2438,6 +2559,7 @@ def create_vector_search_tool( skip: int | Param = 0, filter: FilterExpression | None = None, result_mapper: Callable[[SearchResponse[ModelT]], str | Content | Sequence[Content]] | None = None, + embeddings_options: Mapping[str, Any] | None = None, ) -> FunctionTool: """Create an agent-usable tool backed by vector search. @@ -2454,6 +2576,9 @@ def create_vector_search_tool( its leaf for an absent or null argument. Remaining group children still apply; empty groups are removed recursively. See ``FilterGroup`` for details. result_mapper: Maps each search response to text or one or more multimodal content items. + embeddings_options: Provider options for local query embedding. The vector field + supplies dimensions; a conflicting ``dimensions`` value raises an error. + Requires the search implementation to have an embedding generator. Returns: A function tool with ``query`` and any parameters discovered in ``filter``, ``top``, or ``skip``. @@ -2472,6 +2597,7 @@ def create_vector_search_tool( if isinstance(skip, int): _validate_paging(top=1, skip=skip) + configured_embeddings_options = deepcopy(dict(embeddings_options)) if embeddings_options is not None else None map_result = result_mapper or _default_search_result_mapper configured_filter = snapshot_filter(filter) if filter is not None else None input_schema, param_definitions = _create_search_tool_input_schema( @@ -2519,13 +2645,23 @@ async def search_tool(**arguments: Any) -> list[Content]: resolved_filter, field_names=definition.names if isinstance(definition, VectorStoreCollectionDefinition) else None, ) - results = await search.search( - query, - search_type=search_type, - filter=resolved_filter, - top=invocation_top, - skip=invocation_skip, - ) + if configured_embeddings_options is None: + results = await search.search( + query, + search_type=search_type, + filter=resolved_filter, + top=invocation_top, + skip=invocation_skip, + ) + else: + results = await search.search( + query, + embeddings_options=deepcopy(configured_embeddings_options), + search_type=search_type, + filter=resolved_filter, + top=invocation_top, + skip=invocation_skip, + ) mapped_results: list[Content] = [] consumed_results = 0 async for result in results: @@ -3144,6 +3280,9 @@ def __init__( include_get_tool: bool = True, include_delete_tool: bool = True, include_search_tool: bool = True, + upsert_embeddings_options: Mapping[str, Any] | None = None, + upsert_embeddings_options_by_field: Mapping[str, Mapping[str, Any]] | None = None, + search_embeddings_options: Mapping[str, Any] | None = None, approval_mode: ( Literal["always_require", "never_require"] | Mapping[ @@ -3171,6 +3310,12 @@ def __init__( include_get_tool: Whether to add the default get-by-key tool. include_delete_tool: Whether to add the default delete-by-key tool. include_search_tool: Whether to add the default vector search tool. + upsert_embeddings_options: Provider options for every vector field generated + by the default upsert tool. + upsert_embeddings_options_by_field: Options keyed by logical vector field + name for the default upsert tool. Mutually exclusive with flat options. + search_embeddings_options: Provider options for the default search tool's + local query embedding. approval_mode: One mode for every generated tool, or per-tool overrides merged over the safe defaults. additional_search_tools: Additional caller-configured search tools. @@ -3193,6 +3338,13 @@ def __init__( if not isinstance(value, bool): raise TypeError(f"{name} must be a boolean.") + if not include_upsert_tool and ( + upsert_embeddings_options is not None or upsert_embeddings_options_by_field is not None + ): + raise ValueError("Upsert embedding options require include_upsert_tool=True.") + if not include_search_tool and search_embeddings_options is not None: + raise ValueError("Search embedding options require include_search_tool=True.") + _validate_vector_tool_max_batch_size(max_tool_batch_size) configured_scope_filter = _prepare_vector_tool_filter(collection, scope_filter) approval_modes = self._resolve_approval_modes(approval_mode) @@ -3202,6 +3354,8 @@ def __init__( create_upsert_tool( collection, approval_mode=approval_modes["upsert"], + embeddings_options=upsert_embeddings_options, + embeddings_options_by_field=upsert_embeddings_options_by_field, filter=configured_scope_filter, max_batch_size=max_tool_batch_size, ) @@ -3231,6 +3385,7 @@ def __init__( create_vector_search_tool( cast(SupportsVectorSearch[ModelT], collection), approval_mode=approval_modes["search"], + embeddings_options=search_embeddings_options, filter=configured_scope_filter, ) ) diff --git a/python/packages/core/agent_framework/gemini/__init__.py b/python/packages/core/agent_framework/gemini/__init__.py index 43673773599..7f2cd84eaac 100644 --- a/python/packages/core/agent_framework/gemini/__init__.py +++ b/python/packages/core/agent_framework/gemini/__init__.py @@ -11,9 +11,11 @@ _IMPORTS: dict[str, tuple[str, str]] = { "GeminiChatClient": ("agent_framework_gemini", "agent-framework-gemini"), "GeminiChatOptions": ("agent_framework_gemini", "agent-framework-gemini"), - "GeminiSettings": ("agent_framework_gemini", "agent-framework-gemini"), + "GeminiEmbeddingClient": ("agent_framework_gemini", "agent-framework-gemini"), + "GeminiEmbeddingOptions": ("agent_framework_gemini", "agent-framework-gemini"), "GoogleGeminiSettings": ("agent_framework_gemini", "agent-framework-gemini"), "RawGeminiChatClient": ("agent_framework_gemini", "agent-framework-gemini"), + "RawGeminiEmbeddingClient": ("agent_framework_gemini", "agent-framework-gemini"), "ThinkingConfig": ("agent_framework_gemini", "agent-framework-gemini"), } diff --git a/python/packages/core/agent_framework/gemini/__init__.pyi b/python/packages/core/agent_framework/gemini/__init__.pyi index c9082c9191a..deb6cfc5ed7 100644 --- a/python/packages/core/agent_framework/gemini/__init__.pyi +++ b/python/packages/core/agent_framework/gemini/__init__.pyi @@ -3,17 +3,21 @@ from agent_framework_gemini import ( GeminiChatClient, GeminiChatOptions, - GeminiSettings, + GeminiEmbeddingClient, + GeminiEmbeddingOptions, GoogleGeminiSettings, RawGeminiChatClient, + RawGeminiEmbeddingClient, ThinkingConfig, ) __all__ = [ "GeminiChatClient", "GeminiChatOptions", - "GeminiSettings", + "GeminiEmbeddingClient", + "GeminiEmbeddingOptions", "GoogleGeminiSettings", "RawGeminiChatClient", + "RawGeminiEmbeddingClient", "ThinkingConfig", ] diff --git a/python/packages/core/tests/core/test_gemini_namespace.py b/python/packages/core/tests/core/test_gemini_namespace.py index a59c7bd0e3c..a82b589c95f 100644 --- a/python/packages/core/tests/core/test_gemini_namespace.py +++ b/python/packages/core/tests/core/test_gemini_namespace.py @@ -13,12 +13,16 @@ def test_gemini_namespace_dir_lists_lazy_exports() -> None: for expected in ( "GeminiChatClient", "GeminiChatOptions", - "GeminiSettings", + "GeminiEmbeddingClient", + "GeminiEmbeddingOptions", "GoogleGeminiSettings", "RawGeminiChatClient", + "RawGeminiEmbeddingClient", "ThinkingConfig", ): assert expected in names + assert "GeminiSettings" not in names + assert "GeminiEmbeddingSettings" not in names def test_gemini_namespace_lazy_loads_known_attribute(monkeypatch: pytest.MonkeyPatch) -> None: diff --git a/python/packages/core/tests/core/test_vectors.py b/python/packages/core/tests/core/test_vectors.py index 353fc71827d..acc86b8451c 100644 --- a/python/packages/core/tests/core/test_vectors.py +++ b/python/packages/core/tests/core/test_vectors.py @@ -737,6 +737,82 @@ async def test_collection_serializes_records_and_generates_vectors() -> None: assert embedding_client.options == {"dimensions": 2} +async def test_upsert_embedding_options_are_merged_with_field_dimensions() -> None: + embedding_client = MockEmbeddingClient() + collection = MockCollection(embedding_generator=embedding_client) + caller_options: dict[str, Any] = { + "task_type": "RETRIEVAL_DOCUMENT", + "extra_parameters": {"provider_flag": True}, + "dimensions": 2, + } + + await collection.upsert([Record("one", "first", "document")], embeddings_options=caller_options) + + assert embedding_client.options == caller_options + assert embedding_client.values == ["document"] + cast(dict[str, Any], embedding_client.options)["extra_parameters"]["provider_flag"] = False + assert caller_options["extra_parameters"] == {"provider_flag": True} + assert collection.records["one"]["vector"] == [8.0, 0.5] + + embedding_client.values.clear() + with pytest.raises(ValueError, match="must match its declared dimensions"): + await collection.upsert([Record("two", "second", "document")], embeddings_options={"dimensions": 3}) + assert embedding_client.values == [] + assert "two" not in collection.records + + +@pytest.mark.parametrize("invalid_dimensions", [True, 2.0, None]) +async def test_upsert_rejects_invalid_dimensions_even_if_equal(invalid_dimensions: Any) -> None: + embedding_client = MockEmbeddingClient() + collection = MockCollection(embedding_generator=embedding_client) + with pytest.raises(ValueError, match="must match its declared dimensions"): + await collection.upsert( + [Record("one", "body", "document")], embeddings_options={"dimensions": invalid_dimensions} + ) + assert embedding_client.values == [] + + +async def test_upsert_field_options_are_isolated_and_prevalidated() -> None: + first_client = MockEmbeddingClient() + second_client = MockEmbeddingClient() + definition = VectorStoreCollectionDefinition([ + VectorStoreField("key", name="id"), + VectorStoreField("vector", name="first", dimensions=2, embedding_generator=first_client), + VectorStoreField("vector", name="second", dimensions=2, embedding_generator=second_client), + ]) + handler = VectorStoreRecordHandler(dict, definition=definition) + record = {"id": "one", "first": "first text", "second": "second text"} + + await handler.serialize(record, embeddings_options={"task_type": "RETRIEVAL_DOCUMENT"}) + assert first_client.options == {"task_type": "RETRIEVAL_DOCUMENT", "dimensions": 2} + assert second_client.options == {"task_type": "RETRIEVAL_DOCUMENT", "dimensions": 2} + + by_field = {"first": {"task_type": "RETRIEVAL_DOCUMENT"}, "second": {"task_type": "RETRIEVAL_QUERY"}} + await handler.serialize(record, embeddings_options_by_field=by_field) + assert first_client.options == {"task_type": "RETRIEVAL_DOCUMENT", "dimensions": 2} + assert second_client.options == {"task_type": "RETRIEVAL_QUERY", "dimensions": 2} + assert by_field == {"first": {"task_type": "RETRIEVAL_DOCUMENT"}, "second": {"task_type": "RETRIEVAL_QUERY"}} + + first_client.values.clear() + second_client.values.clear() + with pytest.raises(ValueError, match="second.*declared dimensions"): + await handler.serialize(record, embeddings_options_by_field={"second": {"dimensions": 3}}) + assert first_client.values == second_client.values == [] + + with pytest.raises(ValueError, match="either embeddings_options or embeddings_options_by_field"): + await handler.serialize(record, embeddings_options={}, embeddings_options_by_field={}) + with pytest.raises(ValueError, match="non-generated vector field"): + await handler.serialize(record, embeddings_options_by_field={"missing": {}}) + with pytest.raises(ValueError, match="non-generated vector field"): + await handler.serialize( + record, + generate_vectors=["first"], + embeddings_options_by_field={"second": {}}, + ) + with pytest.raises(ValueError, match="require records and generated vector fields"): + await handler.serialize(record, generate_vectors=False, embeddings_options={}) + + async def test_collection_empty_upsert_skips_embedding_generation() -> None: embedding_client = MockEmbeddingClient() collection = MockCollection(embedding_generator=embedding_client) @@ -1053,6 +1129,36 @@ async def test_vector_search_generates_query_vector_and_forwards_threshold() -> assert responses[1]["score"] == 0.4 +async def test_vector_search_uses_provider_options_and_declared_dimensions() -> None: + embedding_client = MockEmbeddingClient() + collection = MockCollection(embedding_generator=embedding_client) + options = {"task_type": "RETRIEVAL_QUERY"} + + await collection.search("find this", embeddings_options=options) + + assert embedding_client.options == {"task_type": "RETRIEVAL_QUERY", "dimensions": 2} + assert options == {"task_type": "RETRIEVAL_QUERY"} + assert collection.last_search_vector == [9.0, 0.5] + + embedding_client.values.clear() + with pytest.raises(ValueError, match="must match its declared dimensions"): + await collection.search("find this", embeddings_options={"dimensions": 3}) + assert embedding_client.values == [] + + await collection.search(vector=[1.0, 0.0], embeddings_options={"dimensions": 3, **options}) + assert collection.last_search_vector == [1.0, 0.0] + assert embedding_client.values == [] + + await collection.search("keyword terms", vector=[0.0, 1.0], embeddings_options={"dimensions": 3}) + assert collection.last_search_values == "keyword terms" + assert collection.last_search_vector == [0.0, 1.0] + assert embedding_client.values == [] + + await MockCollection().search(vector=[1.0, 0.0], embeddings_options=options) + with pytest.raises(ValueError, match="requires an embedding generator"): + await MockCollection().search("find this", embeddings_options=options) + + async def test_keyword_hybrid_search_uses_single_search_method() -> None: collection = MockCollection() @@ -1333,6 +1439,44 @@ async def test_create_search_tool_returns_mapped_results() -> None: assert result[0].text == "one:0.9" +@pytest.mark.parametrize("search_type", ["vector", "keyword_hybrid"]) +async def test_create_search_tool_forwards_embedding_options(search_type: SearchType) -> None: + embedding_client = MockEmbeddingClient() + collection = MockCollection(embedding_generator=embedding_client) + options = {"task_type": "RETRIEVAL_QUERY"} + + tool = create_vector_search_tool( + collection, + search_type=search_type, + embeddings_options=options, + top=1, + ) + options["task_type"] = "RETRIEVAL_DOCUMENT" + + await tool(query="find this") + + assert embedding_client.values == ["find this"] + assert embedding_client.options == {"task_type": "RETRIEVAL_QUERY", "dimensions": 2} + assert collection.last_search_values == "find this" + assert collection.last_search_vector == [9.0, 0.5] + assert collection.last_search_type == search_type + assert set(tool.parameters()["properties"]) == {"query"} + + +async def test_create_search_tool_rejects_unusable_embedding_options() -> None: + tool = create_vector_search_tool(MockCollection(), embeddings_options={"task_type": "RETRIEVAL_QUERY"}) + with pytest.raises(ValueError, match="requires an embedding generator"): + await tool(query="find this") + + embedding_client = MockEmbeddingClient() + collection = MockCollection(embedding_generator=embedding_client) + tool = create_vector_search_tool(collection, embeddings_options={"dimensions": 3}) + with pytest.raises(ValueError, match="must match its declared dimensions"): + await tool(query="find this") + assert embedding_client.values == [] + assert collection.last_search_type is None + + async def test_create_search_tool_supports_declared_filter_parameters() -> None: collection = MockCollection() collection.records["one"] = {"record_id": "one", "body": "first", "vector": [1.0, 0.0]} @@ -2595,6 +2739,33 @@ async def test_vector_crud_tools_round_trip_records() -> None: assert await get_tool.invoke(arguments={"keys": ["one"]}, skip_parsing=True) == {"records": []} +async def test_upsert_tool_forwards_provider_embedding_options() -> None: + embedding_client = MockEmbeddingClient() + collection = MockCollection(embedding_generator=embedding_client) + options = {"task_type": "RETRIEVAL_DOCUMENT"} + upsert_tool = create_upsert_tool(collection, embeddings_options=options) + options["task_type"] = "RETRIEVAL_QUERY" + + await upsert_tool.invoke( + arguments={"records": [{"id": "one", "text": "body", "vector": "document"}]}, + skip_parsing=True, + ) + + assert embedding_client.options == {"task_type": "RETRIEVAL_DOCUMENT", "dimensions": 2} + assert collection.records["one"]["vector"] == [8.0, 0.5] + with pytest.raises(ValueError, match="either embeddings_options or embeddings_options_by_field"): + create_upsert_tool(collection, embeddings_options={}, embeddings_options_by_field={}) + + by_field = {"vector": {"task_type": "RETRIEVAL_DOCUMENT"}} + by_field_tool = create_upsert_tool(collection, embeddings_options_by_field=by_field) + by_field["vector"]["task_type"] = "RETRIEVAL_QUERY" + await by_field_tool.invoke( + arguments={"records": [{"id": "two", "text": "body", "vector": "more text"}]}, + skip_parsing=True, + ) + assert embedding_client.options == {"task_type": "RETRIEVAL_DOCUMENT", "dimensions": 2} + + async def test_vector_crud_tools_support_auto_generated_keys_when_model_can_omit_them() -> None: definition = VectorStoreCollectionDefinition( [ @@ -2740,6 +2911,36 @@ def test_vector_collection_context_provider_configures_tools_and_approvals() -> assert all(tool.approval_mode == "always_require" for tool in require_all.tools) +async def test_vector_collection_context_provider_forwards_embedding_options() -> None: + embedding_client = MockEmbeddingClient() + collection = MockCollection(embedding_generator=embedding_client) + provider = VectorCollectionContextProvider( + collection, + scope_filter=None, + upsert_embeddings_options={"task_type": "RETRIEVAL_DOCUMENT"}, + search_embeddings_options={"task_type": "RETRIEVAL_QUERY"}, + ) + tools = {tool.name: tool for tool in provider.tools} + + await tools["upsert"].invoke( + arguments={"records": [{"id": "one", "text": "body", "vector": "document"}]}, + skip_parsing=True, + ) + assert embedding_client.options == {"task_type": "RETRIEVAL_DOCUMENT", "dimensions": 2} + + await tools["search"].invoke(arguments={"query": "find this"}, skip_parsing=True) + assert embedding_client.options == {"task_type": "RETRIEVAL_QUERY", "dimensions": 2} + + with pytest.raises(ValueError, match="Upsert embedding options require include_upsert_tool=True"): + VectorCollectionContextProvider( + collection, scope_filter=None, include_upsert_tool=False, upsert_embeddings_options={} + ) + with pytest.raises(ValueError, match="Search embedding options require include_search_tool=True"): + VectorCollectionContextProvider( + collection, scope_filter=None, include_search_tool=False, search_embeddings_options={} + ) + + async def test_vector_collection_context_provider_adds_attributed_context() -> None: collection = MockCollection(embedding_generator=MockEmbeddingClient()) details_tool = create_vector_search_tool(collection, name="search_details") diff --git a/python/packages/gemini/AGENTS.md b/python/packages/gemini/AGENTS.md index 633194061bf..b936bfea0b3 100644 --- a/python/packages/gemini/AGENTS.md +++ b/python/packages/gemini/AGENTS.md @@ -1,15 +1,23 @@ # Gemini Package (agent-framework-gemini) -Integration with Google's Gemini Developer API and Vertex AI via the `google-genai` SDK. +Integration with Google's Gemini Developer API and Enterprise (Vertex AI) via the `google-genai` SDK. ## Core Classes - **`RawGeminiChatClient`** - Lightweight chat client without any layers, for custom pipeline composition - **`GeminiChatClient`** - Full-featured chat client with function invocation, middleware, and telemetry - **`GeminiChatOptions`** - Options TypedDict for Gemini-specific parameters -- **`GeminiSettings`** - Settings loaded from environment variables -- **`GoogleGeminiSettings`** - SDK-standard `GOOGLE_*` settings loaded from environment variables +- **`GoogleGeminiSettings`** - Shared `GOOGLE_*` environment settings for chat and embeddings - **`ThinkingConfig`** - Configuration for extended thinking +- **`RawGeminiEmbeddingClient`** - Text and multimodal embeddings without telemetry +- **`GeminiEmbeddingClient`** - Text and multimodal embeddings with telemetry (defaults to stable `gemini-embedding-2`) +- **`GeminiEmbeddingOptions`** - Per-call embedding model, dimensions, text task, and document title + +`GeminiEmbeddingClient` supports only `gemini-embedding-2` and `gemini-embedding-2-preview`, +requiring per-call task instructions for text strings. Multimodal Google SDK `Content` or +media `Part` inputs receive no task prefix, even with mixed text-and-media parts. Text-only +SDK content is rejected so callers cannot bypass the task requirement. Enterprise accepts +one content per request; the client splits batches there while keeping input order. ## Gemini-specific Options @@ -34,3 +42,15 @@ from agent_framework.gemini import GeminiChatClient client = GeminiChatClient(model="gemini-2.5-flash") response = await client.get_response([Message(role="user", contents=[Content.from_text("Hello")])]) ``` + +```python +from agent_framework.gemini import GeminiEmbeddingClient + +client = GeminiEmbeddingClient() +try: + result = await client.get_embeddings( + ["A document"], options={"task_type": "RETRIEVAL_DOCUMENT", "dimensions": 768} + ) +finally: + await client.close() +``` diff --git a/python/packages/gemini/README.md b/python/packages/gemini/README.md index ae38406b8b7..06d83796831 100644 --- a/python/packages/gemini/README.md +++ b/python/packages/gemini/README.md @@ -8,7 +8,81 @@ pip install agent-framework-gemini --pre ## Gemini Integration -The Gemini integration enables Microsoft Agent Framework applications to call Google Gemini models with familiar chat abstractions, including streaming, tool/function calling, and structured output. +The Gemini integration uses Google's current `google-genai` SDK for chat and text embeddings +through the Gemini Developer API or Gemini Enterprise Agent Platform (formerly Vertex AI). +Chat supports streaming, tool/function calling, and structured output. + +## Gemini Embeddings + +`GeminiEmbeddingClient` defaults to the stable [Gemini Embedding 2](https://ai.google.dev/gemini-api/docs/embeddings) +model (`gemini-embedding-2`). It also accepts `gemini-embedding-2-preview` via +`GOOGLE_EMBEDDING_MODEL` or `model=`; older embedding models are not supported. +It produces one embedding per input: a text string, a Google SDK media `Part`, or +a `google.genai.types.Content` aggregating text and media. Text strings **require** +a `task_type` in each `get_embeddings` call; the client does not choose one by default. + +```python +from agent_framework.gemini import GeminiEmbeddingClient + +client = GeminiEmbeddingClient() +try: + document = await client.get_embeddings( + ["Agent Framework helps build AI agents."], + options={"task_type": "RETRIEVAL_DOCUMENT", "dimensions": 768, "title": "Agent Framework"}, + ) + query = await client.get_embeddings( + ["How do I build an AI agent?"], + options={"task_type": "RETRIEVAL_QUERY", "dimensions": 768}, + ) +finally: + await client.close() +``` + +The client maps Google's [Embedding 2 task instructions](https://ai.google.dev/gemini-api/docs/embeddings). +Use `RETRIEVAL_DOCUMENT` to index documents; pair it with `RETRIEVAL_QUERY` for search, +or `QUESTION_ANSWERING`, `FACT_VERIFICATION`, or `CODE_RETRIEVAL_QUERY` for those specialized +queries. For `CLASSIFICATION`, `CLUSTERING`, or `SEMANTIC_SIMILARITY`, use the same task on +all inputs; `SEMANTIC_SIMILARITY` is not intended for retrieval. Use the same model and +dimensions for indexing and searching. `title` applies to every text in a call, so embed +documents with different titles separately. Generic vector-store embedding generators do +not infer a task type: configure `GeminiEmbeddingClient` as the collection's +`embedding_generator`, pass `embeddings_options={"task_type": "RETRIEVAL_DOCUMENT"}` +to each text upsert, and set `embeddings_options={"task_type": "RETRIEVAL_QUERY"}` +on `create_vector_search_tool`. Core supplies the selected vector field's +dimensions, rejecting a conflicting value. See the +[agent search example](../../samples/02-agents/providers/gemini/gemini_search_agent.py). + +For image, audio, video, PDF, or combined text-and-media input, pass a Google SDK +`Part` or `Content` with at least one media part. The client preserves its parts without +a task prefix, as [Google recommends for multimodal aggregates](https://ai.google.dev/gemini-api/docs/embeddings): + +```python +from google.genai import types + +with open("picture.png", "rb") as image: + image_part = types.Part.from_bytes(data=image.read(), mime_type="image/png") +mixed = types.Content(parts=[types.Part.from_text(text="A landscape photo"), image_part]) +client = GeminiEmbeddingClient() +try: + result = await client.get_embeddings([mixed], options={"dimensions": 768}) +finally: + await client.close() +``` + +Text-only `Part`/`Content` values are rejected: use a string with `task_type` instead. +If a call mixes text strings with media, its `task_type` applies only to the strings; +media content is never prefixed. On Enterprise, Embedding 2 accepts one content per +request, so the client sends multiple inputs as separate, ordered requests. +See the [image search Agent sample](../../samples/02-agents/providers/gemini/gemini_image_search_agent.py) +for cross-modal text-to-image retrieval. Images are embedded without a text task +and upserted with `generate_vectors=False`; the search tool uses +`RETRIEVAL_QUERY` through its `embeddings_options`. + +For embeddings on Enterprise, use the current SDK setting +`GOOGLE_GENAI_USE_ENTERPRISE=true` (or pass `enterprise=True`) together with +`GOOGLE_CLOUD_PROJECT` and `GOOGLE_CLOUD_LOCATION`. The older +`GOOGLE_GENAI_USE_VERTEXAI=true` / `vertexai=True` setting remains supported. +An injected `google.genai.Client` can also provide either authentication mode. ## Structured Output @@ -20,15 +94,18 @@ The connector supports both `google-genai` authentication modes. ### Gemini Developer API -Obtain an API key from [Google AI Studio](https://aistudio.google.com/apikey) and set either the package-prefixed or SDK-standard environment variable: +Obtain an API key from [Google AI Studio](https://aistudio.google.com/apikey) and set the connector's environment variables: ```bash -export GEMINI_API_KEY="your-api-key" -# or: export GOOGLE_API_KEY="your-api-key" -export GEMINI_MODEL="gemini-2.5-flash-lite" -# or: export GOOGLE_MODEL="gemini-2.5-flash-lite" +export GOOGLE_API_KEY="your-api-key" +export GOOGLE_MODEL="gemini-2.5-flash-lite" ``` +The connector no longer reads `GEMINI_API_KEY`, `GEMINI_MODEL`, or +`GEMINI_EMBEDDING_MODEL`. Rename those variables to their `GOOGLE_*` equivalents, +or pass the API key and model explicitly. An injected `google.genai.Client` retains +the Google SDK's own authentication behavior. + ### Vertex AI Set the standard Vertex AI environment variables used by `google-genai`: @@ -49,3 +126,4 @@ See the [Google Gemini samples](../../samples/02-agents/providers/gemini/) for r - Google Search grounding - Google Maps grounding - Built-in code execution +- Text embeddings for document indexing and query retrieval diff --git a/python/packages/gemini/agent_framework_gemini/__init__.py b/python/packages/gemini/agent_framework_gemini/__init__.py index 7a0d014846b..f30f1889e89 100644 --- a/python/packages/gemini/agent_framework_gemini/__init__.py +++ b/python/packages/gemini/agent_framework_gemini/__init__.py @@ -5,11 +5,15 @@ from ._chat_client import ( GeminiChatClient, GeminiChatOptions, - GeminiSettings, GoogleGeminiSettings, RawGeminiChatClient, ThinkingConfig, ) +from ._embedding_client import ( + GeminiEmbeddingClient, + GeminiEmbeddingOptions, + RawGeminiEmbeddingClient, +) try: __version__ = importlib.metadata.version(__name__) @@ -19,9 +23,11 @@ __all__ = [ "GeminiChatClient", "GeminiChatOptions", - "GeminiSettings", + "GeminiEmbeddingClient", + "GeminiEmbeddingOptions", "GoogleGeminiSettings", "RawGeminiChatClient", + "RawGeminiEmbeddingClient", "ThinkingConfig", "__version__", ] diff --git a/python/packages/gemini/agent_framework_gemini/_chat_client.py b/python/packages/gemini/agent_framework_gemini/_chat_client.py index f12f17650a0..129028619c1 100644 --- a/python/packages/gemini/agent_framework_gemini/_chat_client.py +++ b/python/packages/gemini/agent_framework_gemini/_chat_client.py @@ -31,7 +31,7 @@ validate_tool_mode, ) from agent_framework._settings import SecretString, load_settings -from agent_framework._telemetry import get_user_agent, mark_feature_used +from agent_framework._telemetry import mark_feature_used from agent_framework._types import _get_data_bytes # type: ignore[reportPrivateUsage] from agent_framework.exceptions import ( AgentFrameworkException, @@ -48,6 +48,12 @@ from pydantic import BaseModel from ._feature_usage import FeatureIndex +from ._sdk_client import ( + GoogleGeminiSettings, + create_genai_client, + resolve_service_url, + resolve_vertexai_mode, +) if sys.version_info >= (3, 13): from typing import TypeVar # pragma: no cover @@ -69,7 +75,6 @@ __all__ = [ "GeminiChatClient", "GeminiChatOptions", - "GeminiSettings", "GoogleGeminiSettings", "RawGeminiChatClient", "ThinkingConfig", @@ -189,83 +194,12 @@ class GeminiChatOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], to GeminiChatOptionsT = TypeVar("GeminiChatOptionsT", bound=TypedDict, default="GeminiChatOptions", covariant=True) # type: ignore[valid-type] -class GeminiSettings(TypedDict, total=False): - """Gemini configuration settings loaded from environment or .env files.""" - - api_key: SecretString | None - model: str | None - - -class GoogleGeminiSettings(TypedDict, total=False): - """Google SDK configuration settings loaded from ``GOOGLE_*`` environment variables.""" - - api_key: SecretString | None - model: str | None - genai_use_vertexai: bool | None - cloud_project: str | None - cloud_location: str | None - - # endregion -_GEMINI_API_BASE_URL = "https://generativelanguage.googleapis.com" -_VERTEX_AI_BASE_URL = "https://aiplatform.googleapis.com" _DEFAULT_MAX_THOUGHT_SIGNATURES = 256 -def _resolve_vertexai_mode(client: genai.Client, *, fallback: bool | None = None) -> bool: - """Resolve whether a client targets Vertex AI, preferring the instantiated SDK client state.""" - api_client = getattr(client, "_api_client", None) - vertexai = getattr(api_client, "vertexai", None) - if isinstance(vertexai, bool): - return vertexai - return bool(fallback) - - -def _resolve_service_url(client: genai.Client, *, vertexai: bool) -> str: - """Resolve the base service URL from the instantiated SDK client, with a stable fallback.""" - api_client = getattr(client, "_api_client", None) - http_options = getattr(api_client, "_http_options", None) - base_url = getattr(http_options, "base_url", None) - if isinstance(base_url, str) and base_url: - return base_url.rstrip("/") - return _VERTEX_AI_BASE_URL if vertexai else _GEMINI_API_BASE_URL - - -def _validate_client_auth_configuration( - *, - vertexai: bool | None, - api_key: SecretString | None, - project: str | None, - location: str | None, - credentials: Credentials | None, -) -> None: - """Validate supported auth combinations before instantiating the SDK client.""" - if vertexai is not True: - if api_key is None: - raise ValueError( - "Gemini client requires an API key when Vertex AI is not enabled. " - "Set GOOGLE_API_KEY or GEMINI_API_KEY, or pass api_key explicitly." - ) - return - - if api_key is not None or credentials is not None or (project and location): - return - - if project or location: - raise ValueError( - "Gemini client requires both GOOGLE_CLOUD_PROJECT and GOOGLE_CLOUD_LOCATION " - "when Vertex AI is enabled without an API key." - ) - - raise ValueError( - "Gemini client requires Vertex AI credentials or configuration when Vertex AI is enabled. " - "Provide GOOGLE_API_KEY for Vertex AI express mode, pass credentials, or set " - "GOOGLE_CLOUD_PROJECT and GOOGLE_CLOUD_LOCATION." - ) - - # Keys mapping to a different GenerateContentConfig field name _OPTION_TRANSLATIONS: dict[str, str] = { "max_tokens": "max_output_tokens", @@ -360,10 +294,8 @@ def __init__( """Create a raw Gemini chat client. Args: - api_key: Gemini Developer API key. Falls back to environment settings, preferring - ``GOOGLE_API_KEY`` over ``GEMINI_API_KEY``. - model: Default model identifier. Falls back to environment settings, preferring - ``GOOGLE_MODEL`` over ``GEMINI_MODEL``. + api_key: Gemini Developer API key. Falls back to ``GOOGLE_API_KEY``. + model: Default model identifier. Falls back to ``GOOGLE_MODEL``. vertexai: Whether to use Vertex AI endpoints. Falls back to environment settings, using ``GOOGLE_GENAI_USE_VERTEXAI`` when not passed explicitly. project: Google Cloud project ID for Vertex AI. Falls back to environment settings, @@ -384,14 +316,6 @@ def __init__( """ if max_tracked_thought_signatures < 1: raise ValueError("max_tracked_thought_signatures must be greater than 0.") - settings = load_settings( - GeminiSettings, - env_prefix="GEMINI_", - api_key=api_key, - model=model, - env_file_path=env_file_path, - env_file_encoding=env_file_encoding, - ) google_settings = load_settings( GoogleGeminiSettings, env_prefix="GOOGLE_", @@ -405,45 +329,18 @@ def __init__( ) configured_vertexai = google_settings.get("genai_use_vertexai") - if client: - self._genai_client = client - else: - resolved_key = google_settings.get("api_key") or settings.get("api_key") - resolved_project = google_settings.get("cloud_project") - resolved_location = google_settings.get("cloud_location") - _validate_client_auth_configuration( - vertexai=configured_vertexai, - api_key=resolved_key, - project=resolved_project, - location=resolved_location, - credentials=credentials, - ) - - client_kwargs: dict[str, Any] = { - "http_options": {"headers": {"x-goog-api-client": get_user_agent()}}, - } - if configured_vertexai is not None: - client_kwargs["vertexai"] = configured_vertexai - - if resolved_key is not None and ( - configured_vertexai is not True - or (credentials is None and not (resolved_project and resolved_location)) - ): - client_kwargs["api_key"] = resolved_key.get_secret_value() - - if configured_vertexai is True and resolved_project: - client_kwargs["project"] = resolved_project - - if configured_vertexai is True and resolved_location: - client_kwargs["location"] = resolved_location - if configured_vertexai is True and credentials is not None: - client_kwargs["credentials"] = credentials - - self._genai_client = genai.Client(**client_kwargs) + self._genai_client = create_genai_client( + client=client, + api_key=google_settings.get("api_key"), + vertexai=configured_vertexai, + project=google_settings.get("cloud_project"), + location=google_settings.get("cloud_location"), + credentials=credentials, + ) - self._vertexai = _resolve_vertexai_mode(self._genai_client, fallback=configured_vertexai) - self._service_url = _resolve_service_url(self._genai_client, vertexai=self._vertexai) - self.model = google_settings.get("model") or settings.get("model") + self._vertexai = resolve_vertexai_mode(self._genai_client, fallback=configured_vertexai) + self._service_url = resolve_service_url(self._genai_client, vertexai=self._vertexai) + self.model = google_settings.get("model") self.max_tracked_thought_signatures = max_tracked_thought_signatures self._thought_signature_cache: OrderedDict[str, bytes] = OrderedDict() @@ -649,7 +546,7 @@ def _prepare_request( """ model = options.get("model") or self.model if not model: - raise ValueError("Gemini model is required. Set via model parameter or GEMINI_MODEL environment variable.") + raise ValueError("Gemini model is required. Set via model parameter or GOOGLE_MODEL environment variable.") system_instruction, contents = self._prepare_gemini_messages(messages) if call_instructions := options.get("instructions"): @@ -1410,10 +1307,8 @@ def __init__( """Create a Gemini chat client. Args: - api_key: Gemini Developer API key. Falls back to environment settings, preferring - ``GOOGLE_API_KEY`` over ``GEMINI_API_KEY``. - model: Default model identifier. Falls back to environment settings, preferring - ``GOOGLE_MODEL`` over ``GEMINI_MODEL``. + api_key: Gemini Developer API key. Falls back to ``GOOGLE_API_KEY``. + model: Default model identifier. Falls back to ``GOOGLE_MODEL``. vertexai: Whether to use Vertex AI endpoints. Falls back to ``GOOGLE_GENAI_USE_VERTEXAI``. project: Google Cloud project ID for Vertex AI. Falls back to ``GOOGLE_CLOUD_PROJECT``. location: Vertex AI location. Falls back to ``GOOGLE_CLOUD_LOCATION``. diff --git a/python/packages/gemini/agent_framework_gemini/_embedding_client.py b/python/packages/gemini/agent_framework_gemini/_embedding_client.py new file mode 100644 index 00000000000..756b5b77eb3 --- /dev/null +++ b/python/packages/gemini/agent_framework_gemini/_embedding_client.py @@ -0,0 +1,408 @@ +# Copyright (c) Microsoft. All rights reserved. + +from __future__ import annotations + +import math +import sys +from collections.abc import Sequence +from typing import Any, ClassVar, Generic, Literal + +from agent_framework import ( + BaseEmbeddingClient, + Embedding, + EmbeddingGenerationOptions, + GeneratedEmbeddings, + UsageDetails, + load_settings, +) +from agent_framework._settings import SecretString +from agent_framework._telemetry import mark_feature_used +from agent_framework.exceptions import ( + IntegrationException, + IntegrationInvalidAuthException, + IntegrationInvalidRequestException, + IntegrationInvalidResponseException, +) +from agent_framework.observability import EmbeddingTelemetryLayer +from google import genai +from google.auth.credentials import Credentials +from google.genai import types +from google.genai.errors import APIError as GenAIAPIError +from typing_extensions import TypedDict + +from ._feature_usage import FeatureIndex +from ._sdk_client import ( + GoogleGeminiSettings, + create_genai_client, + resolve_service_url, + resolve_vertexai_mode, +) + +if sys.version_info >= (3, 13): + from typing import TypeVar # pragma: no cover +else: + from typing_extensions import TypeVar # pragma: no cover + +_DEFAULT_EMBEDDING_MODEL = "gemini-embedding-2" +_SUPPORTED_EMBEDDING_MODELS = ("gemini-embedding-2", "gemini-embedding-2-preview") +_QUERY_TASK_PREFIXES = { + "RETRIEVAL_QUERY": "search result", + "QUESTION_ANSWERING": "question answering", + "FACT_VERIFICATION": "fact checking", + "CODE_RETRIEVAL_QUERY": "code retrieval", + "CLASSIFICATION": "classification", + "CLUSTERING": "clustering", + "SEMANTIC_SIMILARITY": "sentence similarity", +} + + +class GeminiEmbeddingOptions(EmbeddingGenerationOptions, total=False): + """Google Gemini-specific embedding options. + + ``task_type`` is required per call for text strings, and is not used for + multimodal ``google.genai.types.Content`` or media ``Part`` inputs. The + client formats text according to the Embedding 2 task instructions at + https://ai.google.dev/gemini-api/docs/embeddings. + """ + + task_type: Literal[ + "RETRIEVAL_DOCUMENT", + "RETRIEVAL_QUERY", + "QUESTION_ANSWERING", + "FACT_VERIFICATION", + "CODE_RETRIEVAL_QUERY", + "CLASSIFICATION", + "CLUSTERING", + "SEMANTIC_SIMILARITY", + ] + title: str + + +GeminiEmbeddingOptionsT = TypeVar( + "GeminiEmbeddingOptionsT", + bound=TypedDict, # type: ignore[valid-type] + default="GeminiEmbeddingOptions", + covariant=True, +) + + +def _validate_embedding_model(model: object) -> str: + if not isinstance(model, str) or not model.strip(): + raise ValueError("model must be a non-empty string") + if model not in _SUPPORTED_EMBEDDING_MODELS: + raise ValueError( + f"Unsupported Gemini embedding model {model!r}; use gemini-embedding-2 or gemini-embedding-2-preview." + ) + return model + + +def _prepare_text_for_embedding(text: str, *, task_type: str, title: str | None) -> str: + if task_type == "RETRIEVAL_DOCUMENT": + return f"title: {title or 'none'} | text: {text}" + return f"task: {_QUERY_TASK_PREFIXES[task_type]} | query: {text}" + + +def _prepare_multimodal_content(value: types.Content | types.Part, *, index: int) -> types.Content: + content = types.Content(parts=[value]) if isinstance(value, types.Part) else value + if not content.parts or not any( + part.inline_data is not None or part.file_data is not None for part in content.parts + ): + raise ValueError(f"Multimodal input at index {index} requires a media part; pass text as str with task_type.") + return content + + +def _wrap_gemini_embedding_error(ex: Exception) -> IntegrationException: + """Translate SDK failures into the framework's integration exception hierarchy.""" + if isinstance(ex, ValueError): + return IntegrationInvalidRequestException(f"Invalid Gemini embeddings request: {ex}", inner_exception=ex) + if isinstance(ex, GenAIAPIError): + code = getattr(ex, "code", None) + if code in (401, 403): + return IntegrationInvalidAuthException(f"Gemini embeddings authentication failed: {ex}", inner_exception=ex) + if isinstance(code, int) and 400 <= code < 500: + return IntegrationInvalidRequestException(f"Invalid Gemini embeddings request: {ex}", inner_exception=ex) + return IntegrationException(f"Gemini embeddings request failed: {ex}", inner_exception=ex) + + +class RawGeminiEmbeddingClient( + BaseEmbeddingClient[str | types.Content | types.Part, list[float], GeminiEmbeddingOptionsT], + Generic[GeminiEmbeddingOptionsT], +): + """Generate text and multimodal embeddings via Gemini Developer API or Vertex AI without telemetry. + + Keyword Args: + model: ``gemini-embedding-2`` (default) or ``gemini-embedding-2-preview``. + Set ``GOOGLE_EMBEDDING_MODEL`` to override the default. + api_key: API key, or ``GOOGLE_API_KEY``. + enterprise: Use Gemini Enterprise Agent Platform, or ``GOOGLE_GENAI_USE_ENTERPRISE``. + vertexai: Legacy alias for ``enterprise``, or ``GOOGLE_GENAI_USE_VERTEXAI``. + project: Enterprise (Vertex AI) project, or ``GOOGLE_CLOUD_PROJECT``. + location: Enterprise (Vertex AI) region, or ``GOOGLE_CLOUD_LOCATION``. + credentials: Google Cloud credentials for Enterprise; the SDK can also use ADC. + client: Preconfigured ``genai.Client``; the caller retains ownership of it. + additional_properties: Extra properties stored on the client instance. + env_file_path: Optional ``.env`` file for settings. + env_file_encoding: Encoding for the ``.env`` file. + """ + + OTEL_PROVIDER_NAME: ClassVar[str] = "gcp.gemini" + + def __init__( + self, + *, + model: str | None = None, + api_key: str | SecretString | None = None, + enterprise: bool | None = None, + vertexai: bool | None = None, + project: str | None = None, + location: str | None = None, + credentials: Credentials | None = None, + client: genai.Client | None = None, + additional_properties: dict[str, Any] | None = None, + env_file_path: str | None = None, + env_file_encoding: str | None = None, + ) -> None: + """Initialize a raw Gemini embedding client.""" + google_settings = load_settings( + GoogleGeminiSettings, + env_prefix="GOOGLE_", + api_key=api_key, + embedding_model=model, + genai_use_enterprise=enterprise, + genai_use_vertexai=vertexai, + cloud_project=project, + cloud_location=location, + env_file_path=env_file_path, + env_file_encoding=env_file_encoding, + ) + configured_model = google_settings.get("embedding_model") + self.model = _validate_embedding_model( + _DEFAULT_EMBEDDING_MODEL if configured_model is None else configured_model + ) + + configured_enterprise = google_settings.get("genai_use_enterprise") + configured_vertexai = google_settings.get("genai_use_vertexai") + self._genai_client = create_genai_client( + client=client, + api_key=google_settings.get("api_key"), + enterprise=configured_enterprise, + vertexai=configured_vertexai, + project=google_settings.get("cloud_project"), + location=google_settings.get("cloud_location"), + credentials=credentials, + ) + self._owns_client = client is None + configured_mode = configured_enterprise if configured_enterprise is not None else configured_vertexai + self._vertexai = resolve_vertexai_mode(self._genai_client, fallback=configured_mode) + self._service_url = resolve_service_url(self._genai_client, vertexai=self._vertexai) + + super().__init__(additional_properties=additional_properties) + + def service_url(self) -> str: + """Return the resolved Gemini Developer API or Vertex AI endpoint.""" + return self._service_url + + async def close(self) -> None: + """Close both transports when this client created its own Google SDK client.""" + if self._owns_client: + try: + await self._genai_client.aio.aclose() + finally: + self._genai_client.close() + + async def get_embeddings( + self, + values: Sequence[str | types.Content | types.Part], + *, + options: GeminiEmbeddingOptionsT | None = None, + ) -> GeneratedEmbeddings[list[float], GeminiEmbeddingOptionsT]: + """Generate one embedding per input, preserving the order of texts and multimodal content. + + Args: + values: Text strings or Google SDK ``Content`` / media ``Part`` values. + A ``Content`` may aggregate text and media into one embedding. + options: Model and dimensions, plus a required task type for text strings. + ``title`` applies only to text using ``RETRIEVAL_DOCUMENT``. Multimodal + content is never task-prefixed; omit ``task_type`` for media-only calls. + + Returns: + Embeddings with token counts when the service reports them. + + Raises: + ValueError: If a text lacks a task type, multimodal content lacks media, + or the model, dimensions, task type, or title is invalid. + IntegrationInvalidAuthException: If credentials are rejected by Google. + IntegrationInvalidRequestException: If Google rejects the request. + IntegrationInvalidResponseException: If Google returns malformed embeddings. + IntegrationException: If the SDK request fails for another reason. + """ + if not values: + return GeneratedEmbeddings([], options=options) + + opts: dict[str, Any] = options or {} # type: ignore[assignment] + model = _validate_embedding_model(opts.get("model", self.model)) + + task_type = opts.get("task_type") + title = opts.get("title") + if task_type is not None and ( + not isinstance(task_type, str) + or (task_type != "RETRIEVAL_DOCUMENT" and task_type not in _QUERY_TASK_PREFIXES) + ): + raise ValueError(f"Unsupported Gemini embedding task_type: {task_type!r}.") + if title is not None and (not isinstance(title, str) or not title.strip()): + raise ValueError("title must be a non-empty string") + if title is not None and task_type != "RETRIEVAL_DOCUMENT": + raise ValueError("title requires task_type='RETRIEVAL_DOCUMENT'") + + dimensions = opts.get("dimensions") + if dimensions is not None and ( + isinstance(dimensions, bool) or not isinstance(dimensions, int) or dimensions < 1 + ): + raise ValueError("dimensions must be a positive integer") + + config = types.EmbedContentConfig(output_dimensionality=dimensions) + contents: list[types.Content] = [] + has_text = False + for index, value in enumerate(values): + if isinstance(value, str): + has_text = True + if task_type is None: + raise ValueError("task_type is required for text embeddings.") + contents.append( + types.Content( + parts=[ + types.Part.from_text( + text=_prepare_text_for_embedding(value, task_type=task_type, title=title) + ) + ] + ) + ) + elif isinstance(value, (types.Content, types.Part)): + contents.append(_prepare_multimodal_content(value, index=index)) + else: + raise ValueError(f"Unsupported embedding input at index {index}: {type(value).__name__}.") + if not has_text and task_type is not None: + raise ValueError("task_type is for text strings; omit it when embedding only multimodal content.") + + batches = [[content] for content in contents] if self._vertexai else [contents] + mark_feature_used(FeatureIndex.GEMINI) + raw_embeddings: list[types.ContentEmbedding] = [] + billable_characters = 0 + has_billable_characters = True + for batch in batches: + try: + response = await self._genai_client.aio.models.embed_content( # pyright: ignore[reportUnknownMemberType] + model=model, + contents=batch, + config=config, + ) + except IntegrationException: + raise + except Exception as ex: + raise _wrap_gemini_embedding_error(ex) from ex + + if not isinstance(response, types.EmbedContentResponse): + raise IntegrationInvalidResponseException("Gemini embeddings returned an invalid response.") + batch_embeddings = response.embeddings + if batch_embeddings is None or len(batch_embeddings) != len(batch): + raise IntegrationInvalidResponseException( + f"Gemini embeddings returned {len(batch_embeddings) if batch_embeddings is not None else 0} " + f"vectors for {len(batch)} inputs." + ) + raw_embeddings.extend(batch_embeddings) + if response.metadata is not None and response.metadata.billable_character_count is not None: + billable_characters += response.metadata.billable_character_count + else: + has_billable_characters = False + + embeddings: list[Embedding[list[float]]] = [] + total_tokens = 0 + has_token_counts = True + for item in raw_embeddings: + if not isinstance(item, types.ContentEmbedding): + raise IntegrationInvalidResponseException("Gemini embeddings returned an invalid embedding.") + vector = item.values + if not vector or any( + isinstance(value, bool) or not isinstance(value, (int, float)) or not math.isfinite(value) + for value in vector + ): + raise IntegrationInvalidResponseException("Gemini embeddings returned an invalid vector.") + if dimensions is not None and len(vector) != dimensions: + raise IntegrationInvalidResponseException( + f"Gemini embeddings returned {len(vector)} dimensions; requested {dimensions}." + ) + + stats = item.statistics + extra: dict[str, Any] = {} + if stats is not None: + if stats.truncated is not None: + extra["truncated"] = stats.truncated + count = stats.token_count + if count is not None and ( + isinstance(count, bool) + or not isinstance(count, (int, float)) + or not math.isfinite(count) + or count < 0 + ): + raise IntegrationInvalidResponseException("Gemini embeddings returned an invalid token count.") + if count is not None and float(count).is_integer(): + total_tokens += int(count) + else: + has_token_counts = False + else: + has_token_counts = False + + embeddings.append(Embedding(vector=list(vector), model=model, additional_properties=extra)) + + usage: UsageDetails | None = None + if has_token_counts: + usage = {"input_token_count": total_tokens, "total_token_count": total_tokens} + metadata: dict[str, Any] = {"billable_character_count": billable_characters} if has_billable_characters else {} + return GeneratedEmbeddings(embeddings, options=options, usage=usage, additional_properties=metadata) + + +class GeminiEmbeddingClient( + EmbeddingTelemetryLayer[str | types.Content | types.Part, list[float], GeminiEmbeddingOptionsT], + RawGeminiEmbeddingClient[GeminiEmbeddingOptionsT], + Generic[GeminiEmbeddingOptionsT], +): + """Gemini Developer API and Enterprise text and multimodal embedding client with telemetry. + + Defaults to stable ``gemini-embedding-2``. Pass ``task_type`` in each text call's + options: ``RETRIEVAL_DOCUMENT`` for indexing or ``RETRIEVAL_QUERY`` for searching. + Media ``Part`` and multimodal ``Content`` inputs require no task type. + """ + + OTEL_PROVIDER_NAME: ClassVar[str] = "gcp.gemini" + + def __init__( + self, + *, + model: str | None = None, + api_key: str | SecretString | None = None, + enterprise: bool | None = None, + vertexai: bool | None = None, + project: str | None = None, + location: str | None = None, + credentials: Credentials | None = None, + client: genai.Client | None = None, + otel_provider_name: str | None = None, + additional_properties: dict[str, Any] | None = None, + env_file_path: str | None = None, + env_file_encoding: str | None = None, + ) -> None: + """Initialize a Gemini embedding client with optional telemetry.""" + super().__init__( + model=model, + api_key=api_key, + enterprise=enterprise, + vertexai=vertexai, + project=project, + location=location, + credentials=credentials, + client=client, + otel_provider_name=otel_provider_name, + additional_properties=additional_properties, + env_file_path=env_file_path, + env_file_encoding=env_file_encoding, + ) diff --git a/python/packages/gemini/agent_framework_gemini/_sdk_client.py b/python/packages/gemini/agent_framework_gemini/_sdk_client.py new file mode 100644 index 00000000000..bb7610f52b7 --- /dev/null +++ b/python/packages/gemini/agent_framework_gemini/_sdk_client.py @@ -0,0 +1,127 @@ +# Copyright (c) Microsoft. All rights reserved. + +from __future__ import annotations + +from typing import Any + +from agent_framework._settings import SecretString +from agent_framework._telemetry import get_user_agent +from google import genai +from google.auth.credentials import Credentials +from typing_extensions import TypedDict + + +class GoogleGeminiSettings(TypedDict, total=False): + """Connector settings loaded from ``GOOGLE_*`` environment variables.""" + + api_key: SecretString | None + model: str | None + embedding_model: str | None + genai_use_enterprise: bool | None + genai_use_vertexai: bool | None + cloud_project: str | None + cloud_location: str | None + + +_GEMINI_API_BASE_URL = "https://generativelanguage.googleapis.com" +_VERTEX_AI_BASE_URL = "https://aiplatform.googleapis.com" + + +def resolve_vertexai_mode(client: genai.Client, *, fallback: bool | None = None) -> bool: + """Resolve whether a client targets Vertex AI, preferring the instantiated SDK client state.""" + api_client = getattr(client, "_api_client", None) + vertexai = getattr(api_client, "vertexai", None) + if isinstance(vertexai, bool): + return vertexai + return bool(fallback) + + +def resolve_service_url(client: genai.Client, *, vertexai: bool) -> str: + """Resolve the base service URL from the instantiated SDK client, with a stable fallback.""" + api_client = getattr(client, "_api_client", None) + http_options = getattr(api_client, "_http_options", None) + base_url = getattr(http_options, "base_url", None) + if isinstance(base_url, str) and base_url: + return base_url.rstrip("/") + return _VERTEX_AI_BASE_URL if vertexai else _GEMINI_API_BASE_URL + + +def _validate_client_auth_configuration( + *, + vertexai: bool | None, + api_key: SecretString | None, + project: str | None, + location: str | None, + credentials: Credentials | None, +) -> None: + """Validate supported auth combinations before instantiating the SDK client.""" + if vertexai is not True: + if api_key is None: + raise ValueError( + "Gemini client requires an API key when Vertex AI is not enabled. " + "Set GOOGLE_API_KEY or pass api_key explicitly." + ) + return + + if api_key is not None or credentials is not None or (project and location): + return + + if project or location: + raise ValueError( + "Gemini client requires both GOOGLE_CLOUD_PROJECT and GOOGLE_CLOUD_LOCATION " + "when Vertex AI is enabled without an API key." + ) + + raise ValueError( + "Gemini client requires Vertex AI credentials or configuration when Vertex AI is enabled. " + "Provide GOOGLE_API_KEY for Vertex AI express mode, pass credentials, or set " + "GOOGLE_CLOUD_PROJECT and GOOGLE_CLOUD_LOCATION." + ) + + +def create_genai_client( + *, + client: genai.Client | None, + api_key: SecretString | None, + vertexai: bool | None, + project: str | None, + location: str | None, + credentials: Credentials | None, + enterprise: bool | None = None, +) -> genai.Client: + """Create an SDK client for either Gemini Developer API or Vertex AI, unless one was supplied.""" + if client is not None: + return client + + if api_key is not None and not api_key.get_secret_value().strip(): + raise ValueError("GOOGLE_API_KEY must not be empty when provided.") + if enterprise is not None and vertexai is not None and enterprise != vertexai: + raise ValueError("GOOGLE_GENAI_USE_ENTERPRISE and GOOGLE_GENAI_USE_VERTEXAI cannot disagree.") + use_enterprise = enterprise if enterprise is not None else vertexai + _validate_client_auth_configuration( + vertexai=use_enterprise, + api_key=api_key, + project=project, + location=location, + credentials=credentials, + ) + client_kwargs: dict[str, Any] = { + "http_options": {"headers": {"x-goog-api-client": get_user_agent()}}, + } + if enterprise is not None: + client_kwargs["enterprise"] = enterprise + elif vertexai is not None: + client_kwargs["vertexai"] = vertexai + + if api_key is not None and (use_enterprise is not True or (credentials is None and not (project and location))): + client_kwargs["api_key"] = api_key.get_secret_value() + + if use_enterprise is True and project: + client_kwargs["project"] = project + + if use_enterprise is True and location: + client_kwargs["location"] = location + if use_enterprise is True and credentials is not None: + client_kwargs["credentials"] = credentials + + return genai.Client(**client_kwargs) diff --git a/python/packages/gemini/pyproject.toml b/python/packages/gemini/pyproject.toml index 8480c4fc597..908724158fd 100644 --- a/python/packages/gemini/pyproject.toml +++ b/python/packages/gemini/pyproject.toml @@ -4,7 +4,7 @@ description = "Google Gemini integration for Microsoft Agent Framework." authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}] readme = "README.md" requires-python = ">=3.10" -version = "1.0.0b260910" +version = "1.0.0b260923" license-files = ["LICENSE"] urls.homepage = "https://aka.ms/agent-framework" urls.source = "https://github.com/microsoft/agent-framework/tree/main/python" @@ -25,7 +25,7 @@ classifiers = [ ] dependencies = [ "agent-framework-core>=1.17.0,<2", - "google-genai>=1.69.0,<3.0.0", + "google-genai>=2.25.0,<3.0.0", ] [tool.uv] diff --git a/python/packages/gemini/tests/gemini/test_gemini_embedding_client.py b/python/packages/gemini/tests/gemini/test_gemini_embedding_client.py new file mode 100644 index 00000000000..d50e829bd86 --- /dev/null +++ b/python/packages/gemini/tests/gemini/test_gemini_embedding_client.py @@ -0,0 +1,696 @@ +# Copyright (c) Microsoft. All rights reserved. + +from __future__ import annotations + +import json +import os +from dataclasses import dataclass +from typing import Annotated, Any, cast +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from agent_framework import ( + GeneratedEmbeddings, + InMemoryCollection, + SupportsGetEmbeddings, + VectorStoreField, + create_vector_search_tool, + vectorstoremodel, +) +from agent_framework._settings import SecretString +from agent_framework.exceptions import ( + IntegrationException, + IntegrationInvalidAuthException, + IntegrationInvalidRequestException, + IntegrationInvalidResponseException, +) +from google import genai +from google.genai import errors as genai_errors +from google.genai import types + +from agent_framework_gemini import ( + GeminiEmbeddingClient, + GeminiEmbeddingOptions, + RawGeminiEmbeddingClient, +) +from agent_framework_gemini._feature_usage import FeatureIndex + + +@pytest.fixture +def clear_google_env(monkeypatch: pytest.MonkeyPatch) -> None: + for key in ( + "GEMINI_API_KEY", + "GOOGLE_API_KEY", + "GEMINI_EMBEDDING_MODEL", + "GOOGLE_EMBEDDING_MODEL", + "GOOGLE_GENAI_USE_ENTERPRISE", + "GOOGLE_GENAI_USE_VERTEXAI", + "GOOGLE_CLOUD_PROJECT", + "GOOGLE_CLOUD_LOCATION", + ): + monkeypatch.delenv(key, raising=False) + + +def _make_sdk_client(*, vertexai: bool = False) -> MagicMock: + sdk = MagicMock() + sdk._api_client.vertexai = vertexai + sdk._api_client._http_options.base_url = ( + "https://aiplatform.googleapis.com/" if vertexai else "https://generativelanguage.googleapis.com/" + ) + sdk.aio.models.embed_content = AsyncMock() + sdk.aio.aclose = AsyncMock() + return sdk + + +def _make_client( + *, + model: str | None = "gemini-embedding-2", + vertexai: bool = False, + **kwargs: Any, +) -> tuple[GeminiEmbeddingClient, MagicMock]: + sdk = _make_sdk_client(vertexai=vertexai) + return GeminiEmbeddingClient(client=sdk, model=model, **kwargs), sdk + + +def _response( + vectors: list[list[float]], + *, + counts: list[float] | None = None, + truncated: list[bool] | None = None, +) -> types.EmbedContentResponse: + return types.EmbedContentResponse( + embeddings=[ + types.ContentEmbedding( + values=vector, + statistics=types.ContentEmbeddingStatistics( + token_count=counts[index] if counts is not None else None, + truncated=truncated[index] if truncated is not None else None, + ) + if counts is not None or truncated is not None + else None, + ) + for index, vector in enumerate(vectors) + ], + ) + + +def _texts(contents: list[types.Content]) -> list[str]: + result: list[str] = [] + for content in contents: + assert content.parts is not None and len(content.parts) == 1 + assert content.parts[0].text is not None + result.append(content.parts[0].text) + return result + + +def test_embedding_clients_are_exported_from_provider_namespace() -> None: + from agent_framework.gemini import GeminiEmbeddingClient as NamespacedClient + from agent_framework.gemini import RawGeminiEmbeddingClient as NamespacedRawClient + + assert NamespacedClient is GeminiEmbeddingClient + assert NamespacedRawClient is RawGeminiEmbeddingClient + client, _ = _make_client() + assert isinstance(client, SupportsGetEmbeddings) + + +@pytest.mark.parametrize("api_key", ["explicit-key", SecretString("explicit-key")], ids=["str", "secret"]) +def test_explicit_api_key_unwrapped(api_key: str | SecretString, clear_google_env: None) -> None: + sdk = _make_sdk_client() + with patch("agent_framework_gemini._sdk_client.genai.Client", return_value=sdk) as factory: + client = GeminiEmbeddingClient(api_key=api_key) + + assert factory.call_args.kwargs["api_key"] == "explicit-key" + assert type(factory.call_args.kwargs["api_key"]) is str + assert "x-goog-api-client" in factory.call_args.kwargs["http_options"]["headers"] + assert client.service_url() == "https://generativelanguage.googleapis.com" + + +def test_embedding_uses_google_settings_only(monkeypatch: pytest.MonkeyPatch, clear_google_env: None) -> None: + monkeypatch.setenv("GEMINI_API_KEY", "gemini-key") + monkeypatch.setenv("GOOGLE_API_KEY", "google-key") + monkeypatch.setenv("GEMINI_EMBEDDING_MODEL", "gemini-embedding-001") + monkeypatch.setenv("GOOGLE_EMBEDDING_MODEL", "gemini-embedding-2") + monkeypatch.setenv("GOOGLE_MODEL", "gemini-2.5-flash") + sdk = _make_sdk_client() + with patch("agent_framework_gemini._sdk_client.genai.Client", return_value=sdk) as factory: + client = GeminiEmbeddingClient() + + assert factory.call_args.kwargs["api_key"] == "google-key" + assert client.model == "gemini-embedding-2" + assert GeminiEmbeddingClient(model="gemini-embedding-2-preview", client=sdk).model == "gemini-embedding-2-preview" + + +def test_gemini_only_environment_is_not_supported(monkeypatch: pytest.MonkeyPatch, clear_google_env: None) -> None: + monkeypatch.setenv("GEMINI_API_KEY", "legacy-key") + monkeypatch.setenv("GEMINI_EMBEDDING_MODEL", "gemini-embedding-001") + + with pytest.raises(ValueError, match="GOOGLE_API_KEY"): + GeminiEmbeddingClient() + + injected = GeminiEmbeddingClient(client=_make_sdk_client()) + assert injected.model == "gemini-embedding-2" + + +def test_gemini_embedding_2_is_default(clear_google_env: None) -> None: + client, _ = _make_client(model=None) + assert client.model == "gemini-embedding-2" + + +def test_blank_model_setting_rejected(monkeypatch: pytest.MonkeyPatch, clear_google_env: None) -> None: + monkeypatch.setenv("GOOGLE_EMBEDDING_MODEL", " ") + with pytest.raises(ValueError, match="model must be a non-empty string"): + GeminiEmbeddingClient(client=_make_sdk_client()) + + +def test_vertex_ai_settings_reuse_chat_auth(monkeypatch: pytest.MonkeyPatch, clear_google_env: None) -> None: + monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "true") + monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "project") + monkeypatch.setenv("GOOGLE_CLOUD_LOCATION", "global") + sdk = _make_sdk_client(vertexai=True) + with patch("agent_framework_gemini._sdk_client.genai.Client", return_value=sdk) as factory: + client = GeminiEmbeddingClient() + + assert factory.call_args.kwargs["vertexai"] is True + assert factory.call_args.kwargs["project"] == "project" + assert factory.call_args.kwargs["location"] == "global" + assert "api_key" not in factory.call_args.kwargs + assert client.service_url() == "https://aiplatform.googleapis.com" + + +def test_enterprise_settings_use_current_sdk_mode(monkeypatch: pytest.MonkeyPatch, clear_google_env: None) -> None: + monkeypatch.setenv("GOOGLE_GENAI_USE_ENTERPRISE", "true") + monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "project") + monkeypatch.setenv("GOOGLE_CLOUD_LOCATION", "global") + sdk = _make_sdk_client(vertexai=True) + with patch("agent_framework_gemini._sdk_client.genai.Client", return_value=sdk) as factory: + client = GeminiEmbeddingClient() + + assert factory.call_args.kwargs["enterprise"] is True + assert "vertexai" not in factory.call_args.kwargs + assert factory.call_args.kwargs["project"] == "project" + assert factory.call_args.kwargs["location"] == "global" + assert "api_key" not in factory.call_args.kwargs + assert client.service_url() == "https://aiplatform.googleapis.com" + + +def test_conflicting_enterprise_flags_raise(monkeypatch: pytest.MonkeyPatch, clear_google_env: None) -> None: + monkeypatch.setenv("GOOGLE_API_KEY", "test-key") + monkeypatch.setenv("GOOGLE_GENAI_USE_ENTERPRISE", "true") + monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "false") + with pytest.raises(ValueError, match="cannot disagree"): + GeminiEmbeddingClient() + + +def test_missing_auth_and_incomplete_vertex_config_raise( + monkeypatch: pytest.MonkeyPatch, clear_google_env: None +) -> None: + with pytest.raises(ValueError, match="requires an API key"): + GeminiEmbeddingClient() + + monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "true") + monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "project") + with pytest.raises(ValueError, match="requires both GOOGLE_CLOUD_PROJECT and GOOGLE_CLOUD_LOCATION"): + GeminiEmbeddingClient() + + +def test_injected_vertex_client_controls_service_url() -> None: + client = GeminiEmbeddingClient(client=_make_sdk_client(vertexai=True), vertexai=False) + assert client.service_url() == "https://aiplatform.googleapis.com" + + +async def test_close_only_closes_owned_sdk_client(clear_google_env: None) -> None: + owned_sdk = _make_sdk_client() + with patch("agent_framework_gemini._sdk_client.genai.Client", return_value=owned_sdk): + owned = GeminiEmbeddingClient(api_key="test-key") + await owned.close() + owned_sdk.aio.aclose.assert_awaited_once() + owned_sdk.close.assert_called_once() + + injected, sdk = _make_client() + await injected.close() + sdk.aio.aclose.assert_not_awaited() + sdk.close.assert_not_called() + + +async def test_batch_result_options_usage_and_metadata() -> None: + client, sdk = _make_client() + response = _response([[0.1, 0.2], [0.3, 0.4]], counts=[3.0, 4.0], truncated=[False, True]) + response.metadata = types.EmbedContentMetadata(billable_character_count=16) + sdk.aio.models.embed_content.return_value = response + options: GeminiEmbeddingOptions = { + "model": "gemini-embedding-2-preview", + "dimensions": 2, + "task_type": "RETRIEVAL_DOCUMENT", + "title": "Document title", + } + + with patch("agent_framework_gemini._embedding_client.mark_feature_used") as mark: + result = await client.get_embeddings(["first", "second"], options=options) + + mark.assert_called_once_with(FeatureIndex.GEMINI) + assert isinstance(result, GeneratedEmbeddings) + assert result.options is options + assert [embedding.vector for embedding in result] == [[0.1, 0.2], [0.3, 0.4]] + assert [embedding.dimensions for embedding in result] == [2, 2] + assert [embedding.model for embedding in result] == ["gemini-embedding-2-preview"] * 2 + assert [embedding.additional_properties["truncated"] for embedding in result] == [False, True] + assert result.usage == {"input_token_count": 7, "total_token_count": 7} + assert result.additional_properties == {"billable_character_count": 16} + request = sdk.aio.models.embed_content.call_args.kwargs + assert request["model"] == "gemini-embedding-2-preview" + assert _texts(request["contents"]) == [ + "title: Document title | text: first", + "title: Document title | text: second", + ] + assert request["config"].task_type is None + assert request["config"].title is None + assert request["config"].output_dimensionality == 2 + + +async def test_text_requires_task_type_on_every_call() -> None: + client, sdk = _make_client() + with pytest.raises(ValueError, match="task_type is required for text embeddings"): + await client.get_embeddings(["document"]) + with pytest.raises(ValueError, match="task_type is required for text embeddings"): + await client.get_embeddings(["query"], options={"dimensions": 768}) + sdk.aio.models.embed_content.assert_not_awaited() + + +@pytest.mark.parametrize("mime_type", ["image/png", "audio/mpeg", "video/mp4", "application/pdf"]) +async def test_media_parts_need_no_task_and_are_not_prefixed(mime_type: str) -> None: + client, sdk = _make_client() + sdk.aio.models.embed_content.return_value = _response([[0.1]]) + part = types.Part.from_bytes(data=b"media", mime_type=mime_type) + + result = await client.get_embeddings([part], options={"dimensions": 1}) + + assert result[0].vector == [0.1] + request = sdk.aio.models.embed_content.call_args.kwargs + assert request["contents"][0].parts == [part] + assert request["config"].task_type is None + + +async def test_mixed_batch_preserves_multimodal_content_and_input_order() -> None: + client, sdk = _make_client() + sdk.aio.models.embed_content.return_value = _response([[0.1], [0.2], [0.3]]) + image = types.Part.from_bytes(data=b"image", mime_type="image/png") + aggregate = types.Content(parts=[types.Part.from_text(text="An image of a dog"), image]) + pdf = types.Part.from_uri(file_uri="gs://example/document.pdf", mime_type="application/pdf") + + result = await client.get_embeddings( + ["Find a dog", aggregate, pdf], options={"task_type": "RETRIEVAL_QUERY", "dimensions": 1} + ) + + assert [item.vector for item in result] == [[0.1], [0.2], [0.3]] + contents = sdk.aio.models.embed_content.call_args.kwargs["contents"] + assert _texts(contents[:1]) == ["task: search result | query: Find a dog"] + assert contents[1] is aggregate + assert aggregate.parts is not None and aggregate.parts[0].text == "An image of a dog" + assert contents[2].parts == [pdf] + + +async def test_multimodal_only_rejects_task_type() -> None: + client, sdk = _make_client() + image = types.Part.from_bytes(data=b"image", mime_type="image/png") + with pytest.raises(ValueError, match="omit it when embedding only multimodal content"): + await client.get_embeddings([image], options={"task_type": "RETRIEVAL_QUERY"}) + sdk.aio.models.embed_content.assert_not_awaited() + + +@pytest.mark.parametrize( + "value", + [ + types.Part.from_text(text="text"), + types.Content(parts=[types.Part.from_text(text="text")]), + types.Content(parts=[]), + ], +) +async def test_text_only_sdk_content_must_use_string_with_task(value: types.Content | types.Part) -> None: + client, sdk = _make_client() + with pytest.raises(ValueError, match="requires a media part; pass text as str with task_type"): + await client.get_embeddings([value]) + sdk.aio.models.embed_content.assert_not_awaited() + + +async def test_unsupported_embedding_input_is_rejected() -> None: + client, sdk = _make_client() + with pytest.raises(ValueError, match="Unsupported embedding input at index 0"): + await client.get_embeddings([cast(Any, 123)]) + sdk.aio.models.embed_content.assert_not_awaited() + + +@pytest.mark.parametrize( + ("task_type", "expected"), + [ + ("RETRIEVAL_DOCUMENT", "title: none | text: value"), + ("RETRIEVAL_QUERY", "task: search result | query: value"), + ("QUESTION_ANSWERING", "task: question answering | query: value"), + ("FACT_VERIFICATION", "task: fact checking | query: value"), + ("CODE_RETRIEVAL_QUERY", "task: code retrieval | query: value"), + ("CLASSIFICATION", "task: classification | query: value"), + ("CLUSTERING", "task: clustering | query: value"), + ("SEMANTIC_SIMILARITY", "task: sentence similarity | query: value"), + ], +) +async def test_text_task_instructions(task_type: str, expected: str) -> None: + client, sdk = _make_client() + sdk.aio.models.embed_content.return_value = _response([[0.1]]) + + result = await client.get_embeddings(["value"], options=cast(GeminiEmbeddingOptions, {"task_type": task_type})) + + assert result[0].vector == [0.1] + request = sdk.aio.models.embed_content.call_args.kwargs + assert request["config"].task_type is None + assert _texts(request["contents"]) == [expected] + + +async def test_embedding_2_rejects_unsupported_task_type() -> None: + client, sdk = _make_client() + with pytest.raises(ValueError, match="Unsupported Gemini embedding task_type"): + await client.get_embeddings(["text"], options=cast(GeminiEmbeddingOptions, {"task_type": "UNKNOWN"})) + with pytest.raises(ValueError, match="Unsupported Gemini embedding task_type"): + await client.get_embeddings(["text"], options=cast(GeminiEmbeddingOptions, {"task_type": ""})) + sdk.aio.models.embed_content.assert_not_awaited() + + +@pytest.mark.parametrize("dimensions", [0, -1, True, 1.5, "768"]) +async def test_invalid_dimensions_rejected_before_sdk_call(dimensions: Any) -> None: + client, sdk = _make_client() + with pytest.raises(ValueError, match="dimensions must be a positive integer"): + await client.get_embeddings( + ["text"], options=cast(GeminiEmbeddingOptions, {"task_type": "RETRIEVAL_QUERY", "dimensions": dimensions}) + ) + sdk.aio.models.embed_content.assert_not_awaited() + + +async def test_title_requires_retrieval_document() -> None: + client, sdk = _make_client() + with pytest.raises(ValueError, match="title requires"): + await client.get_embeddings(["text"], options={"title": "name"}) + + document_client, document_sdk = _make_client() + document_sdk.aio.models.embed_content.return_value = _response([[0.1]]) + await document_client.get_embeddings(["text"], options={"task_type": "RETRIEVAL_DOCUMENT", "title": "name"}) + request = document_sdk.aio.models.embed_content.call_args.kwargs + assert request["config"].title is None + assert _texts(request["contents"]) == ["title: name | text: text"] + with pytest.raises(ValueError, match="title must be a non-empty string"): + await document_client.get_embeddings(["text"], options={"task_type": "RETRIEVAL_DOCUMENT", "title": " "}) + + with pytest.raises(ValueError, match="title requires"): + await document_client.get_embeddings(["text"], options={"task_type": "RETRIEVAL_QUERY", "title": "name"}) + sdk.aio.models.embed_content.assert_not_awaited() + + +@pytest.mark.parametrize("model", [123, 0, None, "", " "]) +async def test_model_override_must_be_nonempty_string(model: Any) -> None: + client, sdk = _make_client() + with pytest.raises(ValueError, match="model must be a non-empty string"): + await client.get_embeddings( + ["text"], options=cast(GeminiEmbeddingOptions, {"model": model, "task_type": "RETRIEVAL_QUERY"}) + ) + sdk.aio.models.embed_content.assert_not_awaited() + + +@pytest.mark.parametrize("model", ["gemini-embedding-001", "gemini-embedding-2-other", "models/gemini-embedding-2"]) +async def test_other_embedding_models_are_rejected(model: str) -> None: + with pytest.raises(ValueError, match="use gemini-embedding-2 or gemini-embedding-2-preview"): + GeminiEmbeddingClient(model=model, client=_make_sdk_client()) + + client, sdk = _make_client() + with pytest.raises(ValueError, match="use gemini-embedding-2 or gemini-embedding-2-preview"): + await client.get_embeddings(["text"], options={"model": model, "task_type": "RETRIEVAL_QUERY"}) + sdk.aio.models.embed_content.assert_not_awaited() + + +async def test_empty_input_avoids_sdk_and_feature_mark() -> None: + client, sdk = _make_client(model=None) + with patch("agent_framework_gemini._embedding_client.mark_feature_used") as mark: + result = await client.get_embeddings([]) + assert result == [] + assert result.usage is None + sdk.aio.models.embed_content.assert_not_awaited() + mark.assert_not_called() + + +@pytest.mark.parametrize("vectors", [None, [], [[0.1]], [[0.1], [0.2], [0.3]]]) +async def test_missing_or_wrong_number_of_embeddings_rejected(vectors: list[list[float]] | None) -> None: + client, sdk = _make_client() + sdk.aio.models.embed_content.return_value = ( + types.EmbedContentResponse(embeddings=None) if vectors is None else _response(vectors) + ) + with pytest.raises(IntegrationInvalidResponseException, match="vectors for 2 inputs"): + await client.get_embeddings(["one", "two"], options={"task_type": "RETRIEVAL_QUERY"}) + + +async def test_malformed_response_or_partial_enterprise_batch_rejected() -> None: + client, sdk = _make_client() + sdk.aio.models.embed_content.return_value = None + with pytest.raises(IntegrationInvalidResponseException, match="invalid response"): + await client.get_embeddings(["one"], options={"task_type": "RETRIEVAL_QUERY"}) + + enterprise_client, enterprise_sdk = _make_client(vertexai=True) + enterprise_sdk.aio.models.embed_content.side_effect = [_response([[0.1]]), _response([])] + with pytest.raises(IntegrationInvalidResponseException, match="vectors for 1 inputs"): + await enterprise_client.get_embeddings(["one", "two"], options={"task_type": "RETRIEVAL_QUERY"}) + assert enterprise_sdk.aio.models.embed_content.await_count == 2 + + +@pytest.mark.parametrize("vector", [[], [float("nan")], [float("inf")], ["not a number"], [True]]) +async def test_invalid_vector_rejected(vector: list[Any]) -> None: + client, sdk = _make_client() + sdk.aio.models.embed_content.return_value = types.EmbedContentResponse.model_construct( + embeddings=[types.ContentEmbedding.model_construct(values=vector, statistics=None)] + ) + with pytest.raises(IntegrationInvalidResponseException, match="invalid vector"): + await client.get_embeddings(["text"], options={"task_type": "RETRIEVAL_QUERY"}) + + +async def test_invalid_token_count_rejected() -> None: + client, sdk = _make_client() + sdk.aio.models.embed_content.return_value = types.EmbedContentResponse.model_construct( + embeddings=[ + types.ContentEmbedding.model_construct( + values=[0.1], + statistics=types.ContentEmbeddingStatistics.model_construct(token_count=float("inf")), + ) + ] + ) + with pytest.raises(IntegrationInvalidResponseException, match="invalid token count"): + await client.get_embeddings(["text"], options={"task_type": "RETRIEVAL_QUERY"}) + + +async def test_dimension_mismatch_rejected() -> None: + client, sdk = _make_client() + sdk.aio.models.embed_content.return_value = _response([[0.1]]) + with pytest.raises(IntegrationInvalidResponseException, match="requested 2"): + await client.get_embeddings(["text"], options={"task_type": "RETRIEVAL_QUERY", "dimensions": 2}) + + +async def test_partial_or_fractional_statistics_do_not_report_incorrect_usage() -> None: + client, sdk = _make_client() + sdk.aio.models.embed_content.return_value = _response([[0.1], [0.2]], counts=[3.0, 1.5]) + assert (await client.get_embeddings(["one", "two"], options={"task_type": "RETRIEVAL_QUERY"})).usage is None + + sdk.aio.models.embed_content.return_value = _response([[0.1], [0.2]]) + assert (await client.get_embeddings(["one", "two"], options={"task_type": "RETRIEVAL_QUERY"})).usage is None + + +@pytest.mark.parametrize( + ("error", "expected"), + [ + (genai_errors.ClientError(401, {"error": {"message": "unauthorized"}}), IntegrationInvalidAuthException), + (genai_errors.ClientError(403, {"error": {"message": "forbidden"}}), IntegrationInvalidAuthException), + (genai_errors.ClientError(400, {"error": {"message": "bad request"}}), IntegrationInvalidRequestException), + (genai_errors.ClientError(429, {"error": {"message": "rate limit"}}), IntegrationInvalidRequestException), + (genai_errors.ServerError(500, {"error": {"message": "server error"}}), IntegrationException), + (ValueError("invalid SDK config"), IntegrationInvalidRequestException), + (RuntimeError("network failure"), IntegrationException), + ], +) +async def test_sdk_errors_are_translated(error: Exception, expected: type[Exception]) -> None: + client, sdk = _make_client() + sdk.aio.models.embed_content.side_effect = error + with pytest.raises(expected, match="Gemini embeddings") as caught: + await client.get_embeddings(["text"], options={"task_type": "RETRIEVAL_QUERY"}) + assert caught.value.__cause__ is error + + +async def test_google_genai_serializes_embedding_request() -> None: + requests: list[dict[str, Any]] = [] + + def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path.endswith(":batchEmbedContents") + body = json.loads(request.content) + requests.append(body) + return httpx.Response( + 200, + json={"embeddings": [{"values": [0.1, 0.2], "statistics": {"tokenCount": 2}} for _ in body["requests"]]}, + ) + + sdk = genai.Client( + api_key="test-key", + http_options=types.HttpOptions( + base_url="https://fake.local/", + async_client_args={"transport": httpx.MockTransport(handler)}, + ), + ) + client = GeminiEmbeddingClient(client=sdk) + try: + result = await client.get_embeddings(["one", "two"], options={"task_type": "RETRIEVAL_QUERY", "dimensions": 2}) + assert [embedding.vector for embedding in result] == [[0.1, 0.2], [0.1, 0.2]] + assert result.usage == {"input_token_count": 4, "total_token_count": 4} + assert len(requests[0]["requests"]) == 2 + assert [item["content"]["parts"][0]["text"] for item in requests[0]["requests"]] == [ + "task: search result | query: one", + "task: search result | query: two", + ] + assert all("taskType" not in item and "title" not in item for item in requests[0]["requests"]) + assert all(item["outputDimensionality"] == 2 for item in requests[0]["requests"]) + finally: + await sdk.aio.aclose() + sdk.close() + + +@pytest.mark.filterwarnings("ignore::agent_framework._feature_stage.ExperimentalWarning") +async def test_vector_search_tool_routes_document_and_query_tasks_to_gemini() -> None: + @vectorstoremodel(collection_name="gemini-embedding-task-test") + @dataclass + class Note: + id: Annotated[str, VectorStoreField("key")] + vector: Annotated[ + str | list[float] | None, + VectorStoreField("vector", dimensions=2, distance_function="cosine_similarity"), + ] = None + + client, sdk = _make_client() + sdk.aio.models.embed_content.side_effect = [_response([[0.5, 0.5]]), _response([[0.5, 0.5]])] + collection: InMemoryCollection[str, Note] = InMemoryCollection(Note, embedding_generator=client) + await collection.ensure_collection_exists() + + await collection.upsert([Note("one", "A note")], embeddings_options={"task_type": "RETRIEVAL_DOCUMENT"}) + search_tool = create_vector_search_tool(collection, embeddings_options={"task_type": "RETRIEVAL_QUERY"}, top=1) + results = await search_tool(query="Find the note") + + assert len(results) == 1 + requests = [call.kwargs for call in sdk.aio.models.embed_content.await_args_list] + assert _texts(requests[0]["contents"]) == ["title: none | text: A note"] + assert _texts(requests[1]["contents"]) == ["task: search result | query: Find the note"] + assert all(request["config"].output_dimensionality == 2 for request in requests) + + +async def test_google_genai_serializes_multimodal_aggregate_without_task_prefix() -> None: + requests: list[dict[str, Any]] = [] + + def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path.endswith(":batchEmbedContents") + body = json.loads(request.content) + requests.append(body) + return httpx.Response(200, json={"embeddings": [{"values": [0.1]} for _ in body["requests"]]}) + + sdk = genai.Client( + api_key="test-key", + http_options=types.HttpOptions( + base_url="https://fake.local/", + async_client_args={"transport": httpx.MockTransport(handler)}, + ), + ) + client = GeminiEmbeddingClient(client=sdk) + image = types.Part.from_bytes(data=b"image", mime_type="image/png") + aggregate = types.Content(parts=[types.Part.from_text(text="An image of a dog"), image]) + try: + result = await client.get_embeddings([aggregate], options={"dimensions": 1}) + assert result[0].vector == [0.1] + request = requests[0]["requests"][0] + assert request["content"]["parts"][0]["text"] == "An image of a dog" + assert request["content"]["parts"][1]["inline_data"] == {"data": "aW1hZ2U=", "mime_type": "image/png"} + assert "taskType" not in request + + mixed_batch = await client.get_embeddings( + ["Find the dog", aggregate], options={"task_type": "RETRIEVAL_QUERY", "dimensions": 1} + ) + assert [item.vector for item in mixed_batch] == [[0.1], [0.1]] + assert requests[1]["requests"][0]["content"]["parts"][0]["text"] == "task: search result | query: Find the dog" + assert requests[1]["requests"][1]["content"]["parts"][0]["text"] == "An image of a dog" + finally: + await sdk.aio.aclose() + sdk.close() + + +async def test_google_genai_enterprise_embeds_texts_separately() -> None: + requests: list[dict[str, Any]] = [] + + def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path.endswith(":embedContent") + requests.append(json.loads(request.content)) + return httpx.Response( + 200, + json={"embedding": {"values": [float(len(requests))]}, "metadata": {"billableCharacterCount": 5}}, + ) + + sdk = genai.Client( + enterprise=True, + api_key="test-key", + http_options=types.HttpOptions( + base_url="https://fake.local/", + async_client_args={"transport": httpx.MockTransport(handler)}, + ), + ) + client = GeminiEmbeddingClient(client=sdk) + try: + result = await client.get_embeddings( + ["first", "second"], options={"task_type": "RETRIEVAL_DOCUMENT", "title": "Heading", "dimensions": 1} + ) + assert [embedding.vector for embedding in result] == [[1.0], [2.0]] + assert result.additional_properties == {"billable_character_count": 10} + assert len(requests) == 2 + assert [item["content"]["parts"][0]["text"] for item in requests] == [ + "title: Heading | text: first", + "title: Heading | text: second", + ] + assert all(item["embedContentConfig"]["outputDimensionality"] == 1 for item in requests) + assert all("taskType" not in item["embedContentConfig"] for item in requests) + + image = types.Part.from_bytes(data=b"image", mime_type="image/png") + mixed = types.Content(parts=[types.Part.from_text(text="Unprefixed caption"), image]) + media_result = await client.get_embeddings([mixed], options={"dimensions": 1}) + assert media_result[0].vector == [3.0] + assert requests[2]["content"]["parts"][0]["text"] == "Unprefixed caption" + assert requests[2]["content"]["parts"][1]["inlineData"] == {"data": "aW1hZ2U=", "mime_type": "image/png"} + + mixed_batch = await client.get_embeddings( + ["Find the image", mixed], options={"task_type": "RETRIEVAL_QUERY", "dimensions": 1} + ) + assert [item.vector for item in mixed_batch] == [[4.0], [5.0]] + assert requests[3]["content"]["parts"][0]["text"] == "task: search result | query: Find the image" + assert requests[4]["content"]["parts"][0]["text"] == "Unprefixed caption" + finally: + await sdk.aio.aclose() + sdk.close() + + +def _integration_configured() -> bool: + if not os.getenv("GOOGLE_EMBEDDING_MODEL"): + return False + if os.getenv("GOOGLE_API_KEY"): + return True + return bool( + (os.getenv("GOOGLE_GENAI_USE_ENTERPRISE") or os.getenv("GOOGLE_GENAI_USE_VERTEXAI") or "").lower() + in {"true", "1", "yes", "on"} + and os.getenv("GOOGLE_CLOUD_PROJECT") + and os.getenv("GOOGLE_CLOUD_LOCATION") + ) + + +@pytest.mark.flaky +@pytest.mark.integration +@pytest.mark.skipif(not _integration_configured(), reason="Set Google credentials and an embedding model to run.") +async def test_gemini_embedding_integration() -> None: + client = GeminiEmbeddingClient() + try: + result = await client.get_embeddings(["What is Agent Framework?"], options={"task_type": "RETRIEVAL_QUERY"}) + assert len(result) == 1 + assert result[0].dimensions is not None and result[0].dimensions > 0 + finally: + await client.close() diff --git a/python/packages/gemini/tests/test_gemini_client.py b/python/packages/gemini/tests/test_gemini_client.py index 95a633430bc..db5b4b8a694 100644 --- a/python/packages/gemini/tests/test_gemini_client.py +++ b/python/packages/gemini/tests/test_gemini_client.py @@ -29,7 +29,7 @@ def _has_gemini_integration_credentials() -> bool: """Return whether integration credentials for either Gemini API or Vertex AI appear to be configured.""" - if os.getenv("GEMINI_API_KEY") or os.getenv("GOOGLE_API_KEY"): + if os.getenv("GOOGLE_API_KEY"): return True if os.getenv("GOOGLE_GENAI_USE_VERTEXAI", "").lower() in {"true", "1", "yes", "on"}: @@ -47,7 +47,7 @@ def _has_gemini_integration_credentials() -> bool: reason="Gemini Developer API or Vertex AI credentials not set; skipping integration tests.", ) -_TEST_MODEL = os.getenv("GOOGLE_MODEL") or os.getenv("GEMINI_MODEL", "gemini-2.5-flash-lite") +_TEST_MODEL = os.getenv("GOOGLE_MODEL") or "gemini-2.5-flash-lite" class _ToolListItem(TypedDict): @@ -223,7 +223,7 @@ def test_model_stored_on_instance() -> None: def test_client_created_from_api_key(monkeypatch: pytest.MonkeyPatch) -> None: """Initialises successfully when the API key is supplied via environment variable.""" - monkeypatch.setenv("GEMINI_API_KEY", "test-key-123") + monkeypatch.setenv("GOOGLE_API_KEY", "test-key-123") client = GeminiChatClient(model="gemini-2.5-flash") assert client.model == "gemini-2.5-flash" @@ -233,8 +233,6 @@ def test_client_created_from_google_api_key_env( monkeypatch: pytest.MonkeyPatch, api_key: str | SecretString | None ) -> None: """Initialises successfully when the SDK-standard Google API key environment variable is set.""" - monkeypatch.delenv("GEMINI_API_KEY", raising=False) - monkeypatch.delenv("GEMINI_MODEL", raising=False) monkeypatch.delenv("GOOGLE_GENAI_USE_VERTEXAI", raising=False) monkeypatch.delenv("GOOGLE_CLOUD_PROJECT", raising=False) monkeypatch.delenv("GOOGLE_CLOUD_LOCATION", raising=False) @@ -258,7 +256,6 @@ def test_client_created_from_google_api_key_env( def test_client_created_from_vertex_ai_env(monkeypatch: pytest.MonkeyPatch) -> None: """Initialises a Vertex AI client when the SDK-standard Vertex AI environment variables are set.""" - monkeypatch.delenv("GEMINI_API_KEY", raising=False) monkeypatch.delenv("GOOGLE_API_KEY", raising=False) monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "true") monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "test-project") @@ -278,8 +275,8 @@ def test_client_created_from_vertex_ai_env(monkeypatch: pytest.MonkeyPatch) -> N assert client.service_url() == "https://aiplatform.googleapis.com" -def test_google_settings_take_precedence_over_gemini_aliases(monkeypatch: pytest.MonkeyPatch) -> None: - """Prefers SDK-standard ``GOOGLE_*`` settings when both env families are present.""" +def test_google_settings_are_used_when_gemini_aliases_are_present(monkeypatch: pytest.MonkeyPatch) -> None: + """Only ``GOOGLE_*`` settings configure the connector.""" monkeypatch.setenv("GEMINI_API_KEY", "gemini-key") monkeypatch.setenv("GEMINI_MODEL", "gemini-model") monkeypatch.setenv("GOOGLE_API_KEY", "google-key") @@ -303,10 +300,40 @@ def test_google_settings_take_precedence_over_gemini_aliases(monkeypatch: pytest assert client.service_url() == "https://aiplatform.googleapis.com" +def test_gemini_only_environment_is_not_supported(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("GOOGLE_API_KEY", raising=False) + monkeypatch.delenv("GOOGLE_MODEL", raising=False) + monkeypatch.delenv("GOOGLE_GENAI_USE_VERTEXAI", raising=False) + monkeypatch.delenv("GOOGLE_CLOUD_PROJECT", raising=False) + monkeypatch.delenv("GOOGLE_CLOUD_LOCATION", raising=False) + monkeypatch.setenv("GEMINI_API_KEY", "legacy-key") + monkeypatch.setenv("GEMINI_MODEL", "legacy-model") + + with pytest.raises(ValueError, match="GOOGLE_API_KEY"): + GeminiChatClient() + + injected, _ = _make_gemini_client(model=None) + assert injected.model is None + + +@pytest.mark.parametrize("google_key", ["", " "]) +def test_blank_google_key_does_not_fall_back_to_gemini_key(monkeypatch: pytest.MonkeyPatch, google_key: str) -> None: + monkeypatch.setenv("GOOGLE_API_KEY", google_key) + monkeypatch.setenv("GEMINI_API_KEY", "legacy-key") + monkeypatch.delenv("GOOGLE_GENAI_USE_VERTEXAI", raising=False) + monkeypatch.delenv("GOOGLE_CLOUD_PROJECT", raising=False) + monkeypatch.delenv("GOOGLE_CLOUD_LOCATION", raising=False) + + with ( + patch("agent_framework_gemini._sdk_client.genai.Client") as factory, + pytest.raises(ValueError, match="GOOGLE_API_KEY"), + ): + GeminiChatClient(model="gemini-2.5-flash") + factory.assert_not_called() + + def test_missing_api_key_raises_when_no_client_injected(monkeypatch: pytest.MonkeyPatch) -> None: - """Raises ValueError at construction when neither Gemini API nor Vertex AI settings are available.""" - monkeypatch.delenv("GEMINI_API_KEY", raising=False) - monkeypatch.delenv("GEMINI_MODEL", raising=False) + """Raises ValueError when the configured Google credentials are missing.""" monkeypatch.delenv("GOOGLE_API_KEY", raising=False) monkeypatch.delenv("GOOGLE_GENAI_USE_VERTEXAI", raising=False) monkeypatch.delenv("GOOGLE_CLOUD_PROJECT", raising=False) @@ -321,8 +348,6 @@ def test_vertex_ai_express_mode_uses_api_key( monkeypatch: pytest.MonkeyPatch, api_key: str | SecretString | None ) -> None: """Passes the API key in Vertex AI express mode when no project/location pair is configured.""" - monkeypatch.delenv("GEMINI_API_KEY", raising=False) - monkeypatch.delenv("GEMINI_MODEL", raising=False) monkeypatch.setenv("GOOGLE_API_KEY", "test-key-123") monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "true") monkeypatch.delenv("GOOGLE_CLOUD_PROJECT", raising=False) @@ -345,7 +370,6 @@ def test_vertex_ai_express_mode_uses_api_key( def test_vertex_ai_requires_configuration(monkeypatch: pytest.MonkeyPatch) -> None: """Raises a deterministic error when Vertex AI is enabled without any auth configuration.""" - monkeypatch.delenv("GEMINI_API_KEY", raising=False) monkeypatch.delenv("GOOGLE_API_KEY", raising=False) monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "true") monkeypatch.delenv("GOOGLE_CLOUD_PROJECT", raising=False) @@ -357,7 +381,6 @@ def test_vertex_ai_requires_configuration(monkeypatch: pytest.MonkeyPatch) -> No def test_vertex_ai_requires_project_and_location_together(monkeypatch: pytest.MonkeyPatch) -> None: """Raises a deterministic error when only one Vertex AI location setting is present.""" - monkeypatch.delenv("GEMINI_API_KEY", raising=False) monkeypatch.delenv("GOOGLE_API_KEY", raising=False) monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "true") monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "test-project") @@ -369,7 +392,6 @@ def test_vertex_ai_requires_project_and_location_together(monkeypatch: pytest.Mo async def test_missing_model_raises_on_get_response(monkeypatch: pytest.MonkeyPatch) -> None: """Raises ValueError at call time when no model is set on the client or in options.""" - monkeypatch.delenv("GEMINI_MODEL", raising=False) monkeypatch.delenv("GOOGLE_MODEL", raising=False) client, mock = _make_gemini_client(model=None) # type: ignore[arg-type] mock.aio.models.generate_content = AsyncMock() diff --git a/python/samples/02-agents/providers/gemini/README.md b/python/samples/02-agents/providers/gemini/README.md index c1687368b8f..b7086530b66 100644 --- a/python/samples/02-agents/providers/gemini/README.md +++ b/python/samples/02-agents/providers/gemini/README.md @@ -11,10 +11,25 @@ This folder contains examples demonstrating how to use Google Gemini models with | [`gemini_with_google_search.py`](gemini_with_google_search.py) | Google Search grounding for up-to-date answers. | | [`gemini_with_google_maps.py`](gemini_with_google_maps.py) | Google Maps grounding for location and mapping information. | | [`gemini_with_code_execution.py`](gemini_with_code_execution.py) | Built-in code execution tool for computing precise answers in a sandboxed environment. | +| [`gemini_embeddings.py`](gemini_embeddings.py) | Per-call document and query text embeddings with stable Gemini Embedding 2. | +| [`gemini_search_agent.py`](gemini_search_agent.py) | Document upsert and `create_vector_search_tool` with distinct per-operation embedding options. | +| [`gemini_image_search_agent.py`](gemini_image_search_agent.py) | Cross-modal image indexing and Agent text-to-image search with query embedding options. | + +Run the image search example with two or more local PNG/JPEG files: + +```bash +uv run samples/02-agents/providers/gemini/gemini_image_search_agent.py \ + --query "Which image shows a dog?" photos/dog.jpg photos/cat.png +``` + +Image embeddings are generated without a task prefix. The search tool uses +`RETRIEVAL_QUERY` for text queries and shares the image index's 768 dimensions. ## Environment Variables -- `GOOGLE_MODEL` or `GEMINI_MODEL`: The Gemini model to use (for example, +- `GOOGLE_MODEL`: The Gemini chat model to use (for example, `gemini-2.5-flash-lite` or `gemini-2.5-pro`) -- For Gemini Developer API: `GEMINI_API_KEY` or `GOOGLE_API_KEY` +- For Gemini Developer API: `GOOGLE_API_KEY` - For Vertex AI: `GOOGLE_GENAI_USE_VERTEXAI=true`, `GOOGLE_CLOUD_PROJECT`, and `GOOGLE_CLOUD_LOCATION` +- For embeddings on Enterprise (Vertex AI): `GOOGLE_GENAI_USE_ENTERPRISE=true` is the current SDK equivalent +- `GOOGLE_EMBEDDING_MODEL`: Optional `gemini-embedding-2-preview` override (defaults to `gemini-embedding-2`) diff --git a/python/samples/02-agents/providers/gemini/gemini_advanced.py b/python/samples/02-agents/providers/gemini/gemini_advanced.py index 1d4e3467c23..0e7952b90d5 100644 --- a/python/samples/02-agents/providers/gemini/gemini_advanced.py +++ b/python/samples/02-agents/providers/gemini/gemini_advanced.py @@ -8,8 +8,7 @@ Allows the model to reason through complex problems before responding. -Requires ``GOOGLE_MODEL`` or ``GEMINI_MODEL`` and either Gemini Developer API credentials -(``GEMINI_API_KEY`` or ``GOOGLE_API_KEY``) or Vertex AI settings +Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Vertex AI settings (``GOOGLE_GENAI_USE_VERTEXAI``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). """ diff --git a/python/samples/02-agents/providers/gemini/gemini_basic.py b/python/samples/02-agents/providers/gemini/gemini_basic.py index 15e4cf32908..780651bb208 100644 --- a/python/samples/02-agents/providers/gemini/gemini_basic.py +++ b/python/samples/02-agents/providers/gemini/gemini_basic.py @@ -8,8 +8,7 @@ Covers both non-streaming and streaming responses. -Requires ``GOOGLE_MODEL`` or ``GEMINI_MODEL`` and either Gemini Developer API credentials -(``GEMINI_API_KEY`` or ``GOOGLE_API_KEY``) or Vertex AI settings +Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Vertex AI settings (``GOOGLE_GENAI_USE_VERTEXAI``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). """ diff --git a/python/samples/02-agents/providers/gemini/gemini_embeddings.py b/python/samples/02-agents/providers/gemini/gemini_embeddings.py new file mode 100644 index 00000000000..753a4674adc --- /dev/null +++ b/python/samples/02-agents/providers/gemini/gemini_embeddings.py @@ -0,0 +1,50 @@ +# Copyright (c) Microsoft. All rights reserved. +# /// script +# requires-python = ">=3.10" +# dependencies = ["agent-framework-gemini"] +# /// + +"""Generate document and query text embeddings with the stable Gemini Embedding 2 model. + +Requires ``GOOGLE_API_KEY`` for the Developer API. Enterprise users +can instead set ``GOOGLE_GENAI_USE_ENTERPRISE``, ``GOOGLE_CLOUD_PROJECT``, and +``GOOGLE_CLOUD_LOCATION``. The optional ``GOOGLE_EMBEDDING_MODEL`` setting overrides +the default model. +""" + +import asyncio + +from agent_framework.gemini import GeminiEmbeddingClient +from dotenv import load_dotenv + +load_dotenv() + + +async def main() -> None: + """Embed a document and a search query for the same vector index.""" + # 1. Choose task instructions for each call, not for the client. + client = GeminiEmbeddingClient() + try: + # 2. Use matching dimensions for stored documents and search queries. + document = await client.get_embeddings( + ["Agent Framework helps build and orchestrate AI agents."], + options={"task_type": "RETRIEVAL_DOCUMENT", "title": "Agent Framework", "dimensions": 768}, + ) + query = await client.get_embeddings( + ["How can I orchestrate AI agents?"], + options={"task_type": "RETRIEVAL_QUERY", "dimensions": 768}, + ) + print(f"Document embedding: {document[0].dimensions} dimensions") + print(f"Query embedding: {query[0].dimensions} dimensions") + finally: + await client.close() + + +if __name__ == "__main__": + asyncio.run(main()) + +""" +Sample output: +Document embedding: 768 dimensions +Query embedding: 768 dimensions +""" diff --git a/python/samples/02-agents/providers/gemini/gemini_image_search_agent.py b/python/samples/02-agents/providers/gemini/gemini_image_search_agent.py new file mode 100644 index 00000000000..83a3d778799 --- /dev/null +++ b/python/samples/02-agents/providers/gemini/gemini_image_search_agent.py @@ -0,0 +1,108 @@ +# Copyright (c) Microsoft. All rights reserved. +# /// script +# requires-python = ">=3.10" +# dependencies = ["agent-framework-gemini"] +# /// + +"""Search local images with a text query using Gemini Embedding 2 and an Agent. + +Requires ``GOOGLE_MODEL`` and ``GOOGLE_API_KEY`` for the Developer API, or the +Enterprise project, location, and credential settings. Pass local PNG or JPEG +image paths and a natural-language ``--query``. The in-memory collection is for +demonstration; use a persistent vector store for production data. +""" + +import argparse +import asyncio +from collections.abc import Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Annotated + +from agent_framework import ( + Agent, + InMemoryCollection, + VectorStoreField, + create_vector_search_tool, + vectorstoremodel, +) +from agent_framework.gemini import GeminiChatClient, GeminiEmbeddingClient +from dotenv import load_dotenv +from google.genai import types + +load_dotenv() + +_EMBEDDING_DIMENSIONS = 768 +_IMAGE_MEDIA_TYPES = {".jpg": "image/jpeg", ".jpeg": "image/jpeg", ".png": "image/png"} + + +@vectorstoremodel(collection_name="gemini-image-search") +@dataclass +class IndexedImage: + image_id: Annotated[str, VectorStoreField("key")] + path: Annotated[str, VectorStoreField("data")] + vector: Annotated[ + list[float] | None, + VectorStoreField("vector", dimensions=_EMBEDDING_DIMENSIONS, distance_function="cosine_similarity"), + ] = None + + +async def main(image_paths: Sequence[Path], query: str) -> None: + """Index image vectors without task instructions, then search by a text query.""" + embeddings = GeminiEmbeddingClient() + collection: InMemoryCollection[str, IndexedImage] = InMemoryCollection( + IndexedImage, + embedding_generator=embeddings, + ) + try: + await collection.ensure_collection_exists() + + # 1. Media inputs have no task prefix. Store their vectors explicitly. + records: list[IndexedImage] = [] + for image_path in image_paths: + mime_type = _IMAGE_MEDIA_TYPES.get(image_path.suffix.lower()) + if mime_type is None: + raise ValueError(f"Expected a PNG or JPEG image: {image_path}") + image_part = types.Part.from_bytes(data=image_path.read_bytes(), mime_type=mime_type) + generated = await embeddings.get_embeddings([image_part], options={"dimensions": _EMBEDDING_DIMENSIONS}) + records.append(IndexedImage(str(image_path), str(image_path), generated[0].vector)) + await collection.upsert(records, generate_vectors=False) + + # 2. The helper supplies the query task; Core adds the field dimensions. + search_images = create_vector_search_tool( + collection, + name="search_images", + description="Find images matching a text description and return their file paths.", + top=3, + result_mapper=lambda result: result["record"].path, + embeddings_options={"task_type": "RETRIEVAL_QUERY"}, + ) + async with Agent( + client=GeminiChatClient(), + name="ImageSearchAssistant", + instructions=( + "Use search_images to find matching images. Return their file paths; do not claim to view them." + ), + tools=[search_images], + ) as agent: + response = await agent.run(query) + print(response.text) + finally: + await embeddings.close() + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="Ask a Gemini agent to search local PNG/JPEG images by text.") + parser.add_argument("images", nargs="+", type=Path, help="Local PNG or JPEG files to index") + parser.add_argument("--query", required=True, help="The image-search question for the agent") + args = parser.parse_args() + asyncio.run(main(args.images, args.query)) + +""" +Example: +uv run samples/02-agents/providers/gemini/gemini_image_search_agent.py \ + --query "Which image shows a dog?" photos/dog.jpg photos/cat.png + +Sample output: +The matching image is photos/dog.jpg. +""" diff --git a/python/samples/02-agents/providers/gemini/gemini_search_agent.py b/python/samples/02-agents/providers/gemini/gemini_search_agent.py new file mode 100644 index 00000000000..8ff698ada9b --- /dev/null +++ b/python/samples/02-agents/providers/gemini/gemini_search_agent.py @@ -0,0 +1,97 @@ +# Copyright (c) Microsoft. All rights reserved. +# /// script +# requires-python = ">=3.10" +# dependencies = ["agent-framework-gemini"] +# /// + +"""Give an agent a vector search tool that explicitly selects document and query embedding tasks. + +Requires ``GOOGLE_MODEL`` and ``GOOGLE_API_KEY`` for the Developer API, or the +Enterprise project, location, and credential settings. The in-memory collection +is for demonstration; use a persistent vector store for production data. +""" + +import asyncio +from dataclasses import dataclass +from typing import Annotated + +from agent_framework import ( + Agent, + InMemoryCollection, + VectorStoreField, + create_vector_search_tool, + vectorstoremodel, +) +from agent_framework.gemini import GeminiChatClient, GeminiEmbeddingClient +from dotenv import load_dotenv + +load_dotenv() + +_EMBEDDING_DIMENSIONS = 768 + + +@vectorstoremodel(collection_name="gemini-reference-notes") +@dataclass +class ReferenceNote: + id: Annotated[str, VectorStoreField("key")] + title: Annotated[str, VectorStoreField("data")] + text: Annotated[str, VectorStoreField("data")] + vector: Annotated[ + str | list[float] | None, + VectorStoreField("vector", dimensions=_EMBEDDING_DIMENSIONS, distance_function="cosine_similarity"), + ] = None + + +async def main() -> None: + """Index documents, then let an agent query them using the correct embedding task.""" + embeddings = GeminiEmbeddingClient() + collection: InMemoryCollection[str, ReferenceNote] = InMemoryCollection( + ReferenceNote, + embedding_generator=embeddings, + ) + try: + await collection.ensure_collection_exists() + sources = [ + ("one", "Agent Framework", "Agent Framework builds and orchestrates AI agents."), + ( + "two", + "Gemini embedding tasks", + "Index documents with RETRIEVAL_DOCUMENT; embed search queries with RETRIEVAL_QUERY.", + ), + ] + + # 1. Each document's title applies only to its own embedding request. + for note_id, title, text in sources: + await collection.upsert( + [ReferenceNote(note_id, title, text, text)], + embeddings_options={"task_type": "RETRIEVAL_DOCUMENT", "title": title}, + ) + + # 2. The search helper supplies the query task; Core adds the field dimensions. + search_documents = create_vector_search_tool( + collection, + name="search_documents", + description="Search indexed reference notes for relevant text.", + top=2, + result_mapper=lambda item: f"{item['record'].title}: {item['record'].text}", + embeddings_options={"task_type": "RETRIEVAL_QUERY"}, + ) + async with Agent( + client=GeminiChatClient(), + name="ReferenceAssistant", + instructions="Use search_documents to answer questions about the reference notes. Cite the note title.", + tools=[search_documents], + ) as agent: + response = await agent.run("Which embedding task should I use to search for documents?") + print(response.text) + finally: + await embeddings.close() + + +if __name__ == "__main__": + asyncio.run(main()) + +""" +Sample output: +Use RETRIEVAL_QUERY for search queries (Gemini embedding tasks). +""" diff --git a/python/samples/02-agents/providers/gemini/gemini_with_code_execution.py b/python/samples/02-agents/providers/gemini/gemini_with_code_execution.py index adf7e0bcb25..5aa749ad1c3 100644 --- a/python/samples/02-agents/providers/gemini/gemini_with_code_execution.py +++ b/python/samples/02-agents/providers/gemini/gemini_with_code_execution.py @@ -8,8 +8,7 @@ Allows the model to write and run code in a sandboxed environment to answer questions. -Requires ``GOOGLE_MODEL`` or ``GEMINI_MODEL`` and either Gemini Developer API credentials -(``GEMINI_API_KEY`` or ``GOOGLE_API_KEY``) or Vertex AI settings +Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Vertex AI settings (``GOOGLE_GENAI_USE_VERTEXAI``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). """ diff --git a/python/samples/02-agents/providers/gemini/gemini_with_google_maps.py b/python/samples/02-agents/providers/gemini/gemini_with_google_maps.py index ebabf12fba3..de5715a8e1c 100644 --- a/python/samples/02-agents/providers/gemini/gemini_with_google_maps.py +++ b/python/samples/02-agents/providers/gemini/gemini_with_google_maps.py @@ -8,8 +8,7 @@ Allows Gemini to retrieve location and mapping information before responding. -Requires ``GOOGLE_MODEL`` or ``GEMINI_MODEL`` and either Gemini Developer API credentials -(``GEMINI_API_KEY`` or ``GOOGLE_API_KEY``) or Vertex AI settings +Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Vertex AI settings (``GOOGLE_GENAI_USE_VERTEXAI``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). """ diff --git a/python/samples/02-agents/providers/gemini/gemini_with_google_search.py b/python/samples/02-agents/providers/gemini/gemini_with_google_search.py index 160a74ece21..72acccbf639 100644 --- a/python/samples/02-agents/providers/gemini/gemini_with_google_search.py +++ b/python/samples/02-agents/providers/gemini/gemini_with_google_search.py @@ -8,8 +8,7 @@ Allows Gemini to retrieve up-to-date information from the web before responding. -Requires ``GOOGLE_MODEL`` or ``GEMINI_MODEL`` and either Gemini Developer API credentials -(``GEMINI_API_KEY`` or ``GOOGLE_API_KEY``) or Vertex AI settings +Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Vertex AI settings (``GOOGLE_GENAI_USE_VERTEXAI``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). """ diff --git a/python/samples/02-agents/vector_stores/README.md b/python/samples/02-agents/vector_stores/README.md index d8b855c3223..76fcfc87d89 100644 --- a/python/samples/02-agents/vector_stores/README.md +++ b/python/samples/02-agents/vector_stores/README.md @@ -25,6 +25,20 @@ also use an in-memory collection with OpenAI for embeddings and agent responses. | [`azure_ai_search.py`](azure_ai_search.py) | Native Azure vector/hybrid search with deterministic vectors and a disposable index. | | [`redis_store.py`](redis_store.py) | Native HASH and JSON storage, vector search, filtering, and lifecycle with a disposable Redis server. | +For locally generated vectors, `upsert(..., embeddings_options={...})` applies +provider options to every generated vector field. Use +`embeddings_options_by_field={"field_name": {...}}` for different options on +different logical vector fields; the two arguments are mutually exclusive. +`search(..., embeddings_options={...})` and +`create_vector_search_tool(..., embeddings_options={...})` supply per-query +options. Core adds each selected vector field's declared `dimensions`, and +rejects a conflicting value before embedding or writing records. Upsert +options require generated vectors and cannot be used with `generate_vectors=False`. +Search options require a local embedding generator unless a precomputed +query vector is supplied, in which case they are ignored. See the +[Gemini search sample](../providers/gemini/gemini_search_agent.py) for +document and query task types on the same field. + The Azure sample requires an authorized Azure AI Search service and `az login`. Set `AZURE_SEARCH_ENDPOINT` to your search service. Running the sample creates a uniquely named index, uploads example documents, and deletes that index during diff --git a/python/samples/README.md b/python/samples/README.md index fdfd871ba7c..d4457d1d995 100644 --- a/python/samples/README.md +++ b/python/samples/README.md @@ -156,6 +156,13 @@ variable. | `agent-framework-foundry` | `FoundryChatClient` | `FOUNDRY_MODEL` | `gpt-4o` | | `agent-framework-foundry` | `FoundryAgent` | `FOUNDRY_AGENT_NAME` | `travel-planner` | | `agent-framework-foundry` | `FoundryAgent` | `FOUNDRY_AGENT_VERSION` | `v1` | +| `agent-framework-gemini` | `GeminiChatClient / GeminiEmbeddingClient` | `GOOGLE_API_KEY` | `your-api-key` | +| `agent-framework-gemini` | `GeminiChatClient` | `GOOGLE_MODEL` | `gemini-2.5-flash-lite` | +| `agent-framework-gemini` | `GeminiEmbeddingClient` | `GOOGLE_EMBEDDING_MODEL` | `gemini-embedding-2` | +| `agent-framework-gemini` | `GeminiEmbeddingClient` | `GOOGLE_GENAI_USE_ENTERPRISE` | `true` | +| `agent-framework-gemini` | `GeminiChatClient / GeminiEmbeddingClient` | `GOOGLE_GENAI_USE_VERTEXAI` | `true` | +| `agent-framework-gemini` | `GeminiChatClient / GeminiEmbeddingClient` | `GOOGLE_CLOUD_PROJECT` | `your-project-id` | +| `agent-framework-gemini` | `GeminiChatClient / GeminiEmbeddingClient` | `GOOGLE_CLOUD_LOCATION` | `global` | | `agent-framework-github-copilot` | `GitHubCopilotAgent` | `GITHUB_COPILOT_CLI_PATH` | `copilot` | | `agent-framework-github-copilot` | `GitHubCopilotAgent` | `GITHUB_COPILOT_MODEL` | `gpt-5` | | `agent-framework-github-copilot` | `GitHubCopilotAgent` | `GITHUB_COPILOT_TIMEOUT` | `60` | diff --git a/python/uv.lock b/python/uv.lock index d4daf1e79a3..830c1614d56 100644 --- a/python/uv.lock +++ b/python/uv.lock @@ -705,7 +705,7 @@ requires-dist = [ [[package]] name = "agent-framework-gemini" -version = "1.0.0b260910" +version = "1.0.0b260923" source = { editable = "packages/gemini" } dependencies = [ { name = "agent-framework-core", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, @@ -715,7 +715,7 @@ dependencies = [ [package.metadata] requires-dist = [ { name = "agent-framework-core", editable = "packages/core" }, - { name = "google-genai", specifier = ">=1.69.0,<3.0.0" }, + { name = "google-genai", specifier = ">=2.25.0,<3.0.0" }, ] [[package]] @@ -2418,7 +2418,7 @@ requests = [ [[package]] name = "google-genai" -version = "2.24.0" +version = "2.25.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anyio", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, @@ -2432,9 +2432,9 @@ dependencies = [ { name = "typing-extensions", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, { name = "websockets", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/35/26/176d04f442e82d60cadc74cfaf81f01017062ce274a790455a57aec48c8f/google_genai-2.24.0.tar.gz", hash = "sha256:814b63d05dca4776c5a322699eb30c565040cf8d5fd77c38714076535fb2e019", size = 705409, upload-time = "2026-09-16T22:38:37.367Z" } +sdist = { url = "https://files.pythonhosted.org/packages/62/0a/a3b7856ca840031d4393dbdd97b67fe811b20061315ed68b67b5c85ca80d/google_genai-2.25.0.tar.gz", hash = "sha256:ab603baa5eee0205926ad0f8d7f93e0400df6d67650e99c33e01ab228ea16ad6", size = 699564, upload-time = "2026-09-22T17:23:01.238Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/84/5b/87913fe392ba7b6402577259289ed3740c577956554a674b191a2971d33e/google_genai-2.24.0-py3-none-any.whl", hash = "sha256:59e5f2d88bc2d6ed04aa1aa14cb5629090cfe9458401745e620b24d1fe9893d7", size = 1142675, upload-time = "2026-09-16T22:38:35.332Z" }, + { url = "https://files.pythonhosted.org/packages/5d/a8/178dbb9d1d6cac721b01592e291146a024bae5ee3224e36569348921dd6c/google_genai-2.25.0-py3-none-any.whl", hash = "sha256:2e8f3a5f76ed00d5ffc3153f6089bc7d3511054eed53ad9cfc6095a7dc59b028", size = 1159631, upload-time = "2026-09-22T17:22:59.291Z" }, ] [[package]] From e47d66357d2f679c78226cb3755749a885231819 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Mon, 28 Sep 2026 16:54:08 +0200 Subject: [PATCH 2/4] Python: support Enterprise chat alongside Gemini embeddings --- python/packages/gemini/README.md | 10 +++-- .../agent_framework_gemini/_chat_client.py | 19 +++++++-- python/packages/gemini/pyproject.toml | 2 +- .../gemini/tests/test_gemini_client.py | 40 +++++++++++++++++++ .../02-agents/providers/gemini/README.md | 5 ++- .../providers/gemini/gemini_advanced.py | 4 +- .../providers/gemini/gemini_basic.py | 4 +- .../gemini/gemini_with_code_execution.py | 4 +- .../gemini/gemini_with_google_maps.py | 4 +- .../gemini/gemini_with_google_search.py | 4 +- python/samples/README.md | 2 +- python/uv.lock | 2 +- 12 files changed, 77 insertions(+), 23 deletions(-) diff --git a/python/packages/gemini/README.md b/python/packages/gemini/README.md index 06d83796831..2a11a2e1bf8 100644 --- a/python/packages/gemini/README.md +++ b/python/packages/gemini/README.md @@ -78,7 +78,7 @@ for cross-modal text-to-image retrieval. Images are embedded without a text task and upserted with `generate_vectors=False`; the search tool uses `RETRIEVAL_QUERY` through its `embeddings_options`. -For embeddings on Enterprise, use the current SDK setting +For chat and embeddings on Enterprise, use the current SDK setting `GOOGLE_GENAI_USE_ENTERPRISE=true` (or pass `enterprise=True`) together with `GOOGLE_CLOUD_PROJECT` and `GOOGLE_CLOUD_LOCATION`. The older `GOOGLE_GENAI_USE_VERTEXAI=true` / `vertexai=True` setting remains supported. @@ -106,17 +106,19 @@ The connector no longer reads `GEMINI_API_KEY`, `GEMINI_MODEL`, or or pass the API key and model explicitly. An injected `google.genai.Client` retains the Google SDK's own authentication behavior. -### Vertex AI +### Gemini Enterprise Agent Platform (Vertex AI) -Set the standard Vertex AI environment variables used by `google-genai`: +Set the standard Enterprise environment variables used by `google-genai`: ```bash -export GOOGLE_GENAI_USE_VERTEXAI=true +export GOOGLE_GENAI_USE_ENTERPRISE=true export GOOGLE_CLOUD_PROJECT="your-project-id" export GOOGLE_CLOUD_LOCATION="global" export GOOGLE_MODEL="gemini-2.5-flash-lite" ``` +The older `GOOGLE_GENAI_USE_VERTEXAI=true` setting remains supported. + ## Examples See the [Google Gemini samples](../../samples/02-agents/providers/gemini/) for runnable end-to-end scripts covering: diff --git a/python/packages/gemini/agent_framework_gemini/_chat_client.py b/python/packages/gemini/agent_framework_gemini/_chat_client.py index 129028619c1..18d0fc8afc5 100644 --- a/python/packages/gemini/agent_framework_gemini/_chat_client.py +++ b/python/packages/gemini/agent_framework_gemini/_chat_client.py @@ -281,6 +281,7 @@ def __init__( *, api_key: str | SecretString | None = None, model: str | None = None, + enterprise: bool | None = None, vertexai: bool | None = None, project: str | None = None, location: str | None = None, @@ -296,8 +297,10 @@ def __init__( Args: api_key: Gemini Developer API key. Falls back to ``GOOGLE_API_KEY``. model: Default model identifier. Falls back to ``GOOGLE_MODEL``. - vertexai: Whether to use Vertex AI endpoints. Falls back to environment settings, - using ``GOOGLE_GENAI_USE_VERTEXAI`` when not passed explicitly. + enterprise: Whether to use Gemini Enterprise Agent Platform. Falls back to + ``GOOGLE_GENAI_USE_ENTERPRISE``. + vertexai: Legacy alias for ``enterprise``. Falls back to + ``GOOGLE_GENAI_USE_VERTEXAI`` when not passed explicitly. project: Google Cloud project ID for Vertex AI. Falls back to environment settings, using ``GOOGLE_CLOUD_PROJECT`` when not passed explicitly. location: Vertex AI location. Falls back to environment settings, preferring @@ -321,6 +324,7 @@ def __init__( env_prefix="GOOGLE_", api_key=api_key, model=model, + genai_use_enterprise=enterprise, genai_use_vertexai=vertexai, cloud_project=project, cloud_location=location, @@ -328,17 +332,20 @@ def __init__( env_file_encoding=env_file_encoding, ) + configured_enterprise = google_settings.get("genai_use_enterprise") configured_vertexai = google_settings.get("genai_use_vertexai") self._genai_client = create_genai_client( client=client, api_key=google_settings.get("api_key"), + enterprise=configured_enterprise, vertexai=configured_vertexai, project=google_settings.get("cloud_project"), location=google_settings.get("cloud_location"), credentials=credentials, ) - self._vertexai = resolve_vertexai_mode(self._genai_client, fallback=configured_vertexai) + configured_mode = configured_enterprise if configured_enterprise is not None else configured_vertexai + self._vertexai = resolve_vertexai_mode(self._genai_client, fallback=configured_mode) self._service_url = resolve_service_url(self._genai_client, vertexai=self._vertexai) self.model = google_settings.get("model") self.max_tracked_thought_signatures = max_tracked_thought_signatures @@ -1292,6 +1299,7 @@ def __init__( *, api_key: str | SecretString | None = None, model: str | None = None, + enterprise: bool | None = None, vertexai: bool | None = None, project: str | None = None, location: str | None = None, @@ -1309,7 +1317,9 @@ def __init__( Args: api_key: Gemini Developer API key. Falls back to ``GOOGLE_API_KEY``. model: Default model identifier. Falls back to ``GOOGLE_MODEL``. - vertexai: Whether to use Vertex AI endpoints. Falls back to ``GOOGLE_GENAI_USE_VERTEXAI``. + enterprise: Whether to use Gemini Enterprise Agent Platform. Falls back to + ``GOOGLE_GENAI_USE_ENTERPRISE``. + vertexai: Legacy alias for ``enterprise``. Falls back to ``GOOGLE_GENAI_USE_VERTEXAI``. project: Google Cloud project ID for Vertex AI. Falls back to ``GOOGLE_CLOUD_PROJECT``. location: Vertex AI location. Falls back to ``GOOGLE_CLOUD_LOCATION``. credentials: Google Cloud credentials for Vertex AI. When omitted, the SDK can use @@ -1329,6 +1339,7 @@ def __init__( super().__init__( api_key=api_key, model=model, + enterprise=enterprise, vertexai=vertexai, project=project, location=location, diff --git a/python/packages/gemini/pyproject.toml b/python/packages/gemini/pyproject.toml index 908724158fd..1e8dd92e844 100644 --- a/python/packages/gemini/pyproject.toml +++ b/python/packages/gemini/pyproject.toml @@ -4,7 +4,7 @@ description = "Google Gemini integration for Microsoft Agent Framework." authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}] readme = "README.md" requires-python = ">=3.10" -version = "1.0.0b260923" +version = "1.0.0b260928" license-files = ["LICENSE"] urls.homepage = "https://aka.ms/agent-framework" urls.source = "https://github.com/microsoft/agent-framework/tree/main/python" diff --git a/python/packages/gemini/tests/test_gemini_client.py b/python/packages/gemini/tests/test_gemini_client.py index db5b4b8a694..a6633817db5 100644 --- a/python/packages/gemini/tests/test_gemini_client.py +++ b/python/packages/gemini/tests/test_gemini_client.py @@ -50,6 +50,12 @@ def _has_gemini_integration_credentials() -> bool: _TEST_MODEL = os.getenv("GOOGLE_MODEL") or "gemini-2.5-flash-lite" +@pytest.fixture(autouse=True) +def clear_enterprise_env(monkeypatch: pytest.MonkeyPatch) -> None: + """Isolate chat tests from an externally configured Enterprise SDK mode.""" + monkeypatch.delenv("GOOGLE_GENAI_USE_ENTERPRISE", raising=False) + + class _ToolListItem(TypedDict): title: str description: NotRequired[str] @@ -275,6 +281,40 @@ def test_client_created_from_vertex_ai_env(monkeypatch: pytest.MonkeyPatch) -> N assert client.service_url() == "https://aiplatform.googleapis.com" +def test_chat_clients_created_from_enterprise_env(monkeypatch: pytest.MonkeyPatch) -> None: + """Both chat layers share the current Enterprise settings with the embedding client.""" + monkeypatch.delenv("GOOGLE_API_KEY", raising=False) + monkeypatch.delenv("GOOGLE_GENAI_USE_VERTEXAI", raising=False) + monkeypatch.setenv("GOOGLE_GENAI_USE_ENTERPRISE", "true") + monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "project") + monkeypatch.setenv("GOOGLE_CLOUD_LOCATION", "global") + + mock_client = MagicMock() + mock_client._api_client.vertexai = True + mock_client._api_client._http_options.base_url = "https://aiplatform.googleapis.com/" + with patch("agent_framework_gemini._sdk_client.genai.Client", return_value=mock_client) as factory: + raw = RawGeminiChatClient() + full = GeminiChatClient() + + assert factory.call_count == 2 + for call in factory.call_args_list: + assert call.kwargs["enterprise"] is True + assert "vertexai" not in call.kwargs + assert call.kwargs["project"] == "project" + assert call.kwargs["location"] == "global" + assert "api_key" not in call.kwargs + assert raw.service_url() == full.service_url() == "https://aiplatform.googleapis.com" + + +def test_chat_rejects_conflicting_enterprise_and_vertex_flags(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("GOOGLE_GENAI_USE_ENTERPRISE", "true") + monkeypatch.setenv("GOOGLE_GENAI_USE_VERTEXAI", "false") + monkeypatch.setenv("GOOGLE_API_KEY", "test-key") + + with pytest.raises(ValueError, match="cannot disagree"): + GeminiChatClient(model="gemini-2.5-flash") + + def test_google_settings_are_used_when_gemini_aliases_are_present(monkeypatch: pytest.MonkeyPatch) -> None: """Only ``GOOGLE_*`` settings configure the connector.""" monkeypatch.setenv("GEMINI_API_KEY", "gemini-key") diff --git a/python/samples/02-agents/providers/gemini/README.md b/python/samples/02-agents/providers/gemini/README.md index b7086530b66..dc93169aacd 100644 --- a/python/samples/02-agents/providers/gemini/README.md +++ b/python/samples/02-agents/providers/gemini/README.md @@ -30,6 +30,7 @@ Image embeddings are generated without a task prefix. The search tool uses - `GOOGLE_MODEL`: The Gemini chat model to use (for example, `gemini-2.5-flash-lite` or `gemini-2.5-pro`) - For Gemini Developer API: `GOOGLE_API_KEY` -- For Vertex AI: `GOOGLE_GENAI_USE_VERTEXAI=true`, `GOOGLE_CLOUD_PROJECT`, and `GOOGLE_CLOUD_LOCATION` -- For embeddings on Enterprise (Vertex AI): `GOOGLE_GENAI_USE_ENTERPRISE=true` is the current SDK equivalent +- For Gemini Enterprise Agent Platform (chat and embeddings): `GOOGLE_GENAI_USE_ENTERPRISE=true`, + `GOOGLE_CLOUD_PROJECT`, and `GOOGLE_CLOUD_LOCATION`. The older + `GOOGLE_GENAI_USE_VERTEXAI=true` setting remains supported. - `GOOGLE_EMBEDDING_MODEL`: Optional `gemini-embedding-2-preview` override (defaults to `gemini-embedding-2`) diff --git a/python/samples/02-agents/providers/gemini/gemini_advanced.py b/python/samples/02-agents/providers/gemini/gemini_advanced.py index 0e7952b90d5..7e544a423a2 100644 --- a/python/samples/02-agents/providers/gemini/gemini_advanced.py +++ b/python/samples/02-agents/providers/gemini/gemini_advanced.py @@ -8,8 +8,8 @@ Allows the model to reason through complex problems before responding. -Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Vertex AI settings -(``GOOGLE_GENAI_USE_VERTEXAI``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). +Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Enterprise settings +(``GOOGLE_GENAI_USE_ENTERPRISE``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). """ import asyncio diff --git a/python/samples/02-agents/providers/gemini/gemini_basic.py b/python/samples/02-agents/providers/gemini/gemini_basic.py index 780651bb208..0569593f51b 100644 --- a/python/samples/02-agents/providers/gemini/gemini_basic.py +++ b/python/samples/02-agents/providers/gemini/gemini_basic.py @@ -8,8 +8,8 @@ Covers both non-streaming and streaming responses. -Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Vertex AI settings -(``GOOGLE_GENAI_USE_VERTEXAI``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). +Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Enterprise settings +(``GOOGLE_GENAI_USE_ENTERPRISE``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). """ import asyncio diff --git a/python/samples/02-agents/providers/gemini/gemini_with_code_execution.py b/python/samples/02-agents/providers/gemini/gemini_with_code_execution.py index 5aa749ad1c3..c38564e5eee 100644 --- a/python/samples/02-agents/providers/gemini/gemini_with_code_execution.py +++ b/python/samples/02-agents/providers/gemini/gemini_with_code_execution.py @@ -8,8 +8,8 @@ Allows the model to write and run code in a sandboxed environment to answer questions. -Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Vertex AI settings -(``GOOGLE_GENAI_USE_VERTEXAI``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). +Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Enterprise settings +(``GOOGLE_GENAI_USE_ENTERPRISE``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). """ import asyncio diff --git a/python/samples/02-agents/providers/gemini/gemini_with_google_maps.py b/python/samples/02-agents/providers/gemini/gemini_with_google_maps.py index de5715a8e1c..71f4833af5f 100644 --- a/python/samples/02-agents/providers/gemini/gemini_with_google_maps.py +++ b/python/samples/02-agents/providers/gemini/gemini_with_google_maps.py @@ -8,8 +8,8 @@ Allows Gemini to retrieve location and mapping information before responding. -Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Vertex AI settings -(``GOOGLE_GENAI_USE_VERTEXAI``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). +Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Enterprise settings +(``GOOGLE_GENAI_USE_ENTERPRISE``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). """ import asyncio diff --git a/python/samples/02-agents/providers/gemini/gemini_with_google_search.py b/python/samples/02-agents/providers/gemini/gemini_with_google_search.py index 72acccbf639..72e5a488ae1 100644 --- a/python/samples/02-agents/providers/gemini/gemini_with_google_search.py +++ b/python/samples/02-agents/providers/gemini/gemini_with_google_search.py @@ -8,8 +8,8 @@ Allows Gemini to retrieve up-to-date information from the web before responding. -Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Vertex AI settings -(``GOOGLE_GENAI_USE_VERTEXAI``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). +Requires ``GOOGLE_MODEL`` and either ``GOOGLE_API_KEY`` for the Developer API or Enterprise settings +(``GOOGLE_GENAI_USE_ENTERPRISE``, ``GOOGLE_CLOUD_PROJECT``, and ``GOOGLE_CLOUD_LOCATION``). """ import asyncio diff --git a/python/samples/README.md b/python/samples/README.md index d4457d1d995..a53b7eed88f 100644 --- a/python/samples/README.md +++ b/python/samples/README.md @@ -159,7 +159,7 @@ variable. | `agent-framework-gemini` | `GeminiChatClient / GeminiEmbeddingClient` | `GOOGLE_API_KEY` | `your-api-key` | | `agent-framework-gemini` | `GeminiChatClient` | `GOOGLE_MODEL` | `gemini-2.5-flash-lite` | | `agent-framework-gemini` | `GeminiEmbeddingClient` | `GOOGLE_EMBEDDING_MODEL` | `gemini-embedding-2` | -| `agent-framework-gemini` | `GeminiEmbeddingClient` | `GOOGLE_GENAI_USE_ENTERPRISE` | `true` | +| `agent-framework-gemini` | `GeminiChatClient / GeminiEmbeddingClient` | `GOOGLE_GENAI_USE_ENTERPRISE` | `true` | | `agent-framework-gemini` | `GeminiChatClient / GeminiEmbeddingClient` | `GOOGLE_GENAI_USE_VERTEXAI` | `true` | | `agent-framework-gemini` | `GeminiChatClient / GeminiEmbeddingClient` | `GOOGLE_CLOUD_PROJECT` | `your-project-id` | | `agent-framework-gemini` | `GeminiChatClient / GeminiEmbeddingClient` | `GOOGLE_CLOUD_LOCATION` | `global` | diff --git a/python/uv.lock b/python/uv.lock index 830c1614d56..bdf83d2ae76 100644 --- a/python/uv.lock +++ b/python/uv.lock @@ -705,7 +705,7 @@ requires-dist = [ [[package]] name = "agent-framework-gemini" -version = "1.0.0b260923" +version = "1.0.0b260928" source = { editable = "packages/gemini" } dependencies = [ { name = "agent-framework-core", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" }, From 57ed68771fcc8dbdb4ac960a7ec416341768ffa8 Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Tue, 29 Sep 2026 08:08:02 +0200 Subject: [PATCH 3/4] Python: exercise default Gemini and Enterprise integration paths --- .../packages/core/tests/core/test_vectors.py | 7 +++-- .../gemini/test_gemini_embedding_client.py | 9 ++++-- .../gemini/tests/test_gemini_client.py | 30 ++++++++++++++++--- 3 files changed, 36 insertions(+), 10 deletions(-) diff --git a/python/packages/core/tests/core/test_vectors.py b/python/packages/core/tests/core/test_vectors.py index acc86b8451c..e5e05331a56 100644 --- a/python/packages/core/tests/core/test_vectors.py +++ b/python/packages/core/tests/core/test_vectors.py @@ -91,7 +91,7 @@ class MockEmbeddingClient(BaseEmbeddingClient): def __init__(self) -> None: super().__init__() self.values: list[Any] = [] - self.options: EmbeddingGenerationOptions | None = None + self.options: dict[str, Any] | None = None async def get_embeddings( self, @@ -100,7 +100,7 @@ async def get_embeddings( options: EmbeddingGenerationOptions | None = None, ) -> GeneratedEmbeddings[list[float]]: self.values = list(values) - self.options = options + self.options = dict(options) if options is not None else None return GeneratedEmbeddings([Embedding(vector=[float(len(str(value))), 0.5]) for value in values]) @@ -750,7 +750,8 @@ async def test_upsert_embedding_options_are_merged_with_field_dimensions() -> No assert embedding_client.options == caller_options assert embedding_client.values == ["document"] - cast(dict[str, Any], embedding_client.options)["extra_parameters"]["provider_flag"] = False + assert embedding_client.options is not None + embedding_client.options["extra_parameters"]["provider_flag"] = False assert caller_options["extra_parameters"] == {"provider_flag": True} assert collection.records["one"]["vector"] == [8.0, 0.5] diff --git a/python/packages/gemini/tests/gemini/test_gemini_embedding_client.py b/python/packages/gemini/tests/gemini/test_gemini_embedding_client.py index d50e829bd86..ff45354dc77 100644 --- a/python/packages/gemini/tests/gemini/test_gemini_embedding_client.py +++ b/python/packages/gemini/tests/gemini/test_gemini_embedding_client.py @@ -671,8 +671,6 @@ def handler(request: httpx.Request) -> httpx.Response: def _integration_configured() -> bool: - if not os.getenv("GOOGLE_EMBEDDING_MODEL"): - return False if os.getenv("GOOGLE_API_KEY"): return True return bool( @@ -683,9 +681,14 @@ def _integration_configured() -> bool: ) +def test_integration_gate_uses_default_embedding_model(monkeypatch: pytest.MonkeyPatch, clear_google_env: None) -> None: + monkeypatch.setenv("GOOGLE_API_KEY", "test-key") + assert _integration_configured() + + @pytest.mark.flaky @pytest.mark.integration -@pytest.mark.skipif(not _integration_configured(), reason="Set Google credentials and an embedding model to run.") +@pytest.mark.skipif(not _integration_configured(), reason="Set GOOGLE_API_KEY or Enterprise credentials to run.") async def test_gemini_embedding_integration() -> None: client = GeminiEmbeddingClient() try: diff --git a/python/packages/gemini/tests/test_gemini_client.py b/python/packages/gemini/tests/test_gemini_client.py index a6633817db5..53748972401 100644 --- a/python/packages/gemini/tests/test_gemini_client.py +++ b/python/packages/gemini/tests/test_gemini_client.py @@ -28,11 +28,12 @@ def _has_gemini_integration_credentials() -> bool: - """Return whether integration credentials for either Gemini API or Vertex AI appear to be configured.""" + """Return whether Developer API or Enterprise credentials appear to be configured.""" if os.getenv("GOOGLE_API_KEY"): return True - if os.getenv("GOOGLE_GENAI_USE_VERTEXAI", "").lower() in {"true", "1", "yes", "on"}: + enterprise_mode = os.getenv("GOOGLE_GENAI_USE_ENTERPRISE") or os.getenv("GOOGLE_GENAI_USE_VERTEXAI") or "" + if enterprise_mode.lower() in {"true", "1", "yes", "on"}: return bool( os.getenv("GOOGLE_CLOUD_PROJECT") or os.getenv("GOOGLE_APPLICATION_CREDENTIALS") @@ -51,9 +52,30 @@ def _has_gemini_integration_credentials() -> bool: @pytest.fixture(autouse=True) -def clear_enterprise_env(monkeypatch: pytest.MonkeyPatch) -> None: - """Isolate chat tests from an externally configured Enterprise SDK mode.""" +def clear_enterprise_env(request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch) -> None: + """Isolate unit tests without removing the configured mode from integration tests.""" + if request.node.get_closest_marker("integration") is None: + monkeypatch.delenv("GOOGLE_GENAI_USE_ENTERPRISE", raising=False) + + +@pytest.mark.parametrize("mode", ["GOOGLE_GENAI_USE_ENTERPRISE", "GOOGLE_GENAI_USE_VERTEXAI"]) +def test_enterprise_integration_gate_accepts_either_mode(monkeypatch: pytest.MonkeyPatch, mode: str) -> None: + monkeypatch.delenv("GOOGLE_API_KEY", raising=False) monkeypatch.delenv("GOOGLE_GENAI_USE_ENTERPRISE", raising=False) + monkeypatch.delenv("GOOGLE_GENAI_USE_VERTEXAI", raising=False) + monkeypatch.setenv("GOOGLE_CLOUD_PROJECT", "test-project") + monkeypatch.setenv(mode, "true") + + assert _has_gemini_integration_credentials() + + +@pytest.mark.integration +@pytest.mark.skipif( + os.getenv("GOOGLE_GENAI_USE_ENTERPRISE", "").lower() not in {"true", "1", "yes", "on"}, + reason="Enterprise mode not configured.", +) +def test_integration_fixture_preserves_enterprise_mode() -> None: + assert os.getenv("GOOGLE_GENAI_USE_ENTERPRISE", "").lower() in {"true", "1", "yes", "on"} class _ToolListItem(TypedDict): From 2e0cb5ea840e507fd245b68373cf316f53aecebc Mon Sep 17 00:00:00 2001 From: eavanvalkenburg Date: Tue, 29 Sep 2026 08:56:54 +0200 Subject: [PATCH 4/4] Python: centralize Gemini SDK routing metadata --- python/packages/gemini/AGENTS.md | 2 + .../agent_framework_gemini/_chat_client.py | 7 +- .../_embedding_client.py | 7 +- .../agent_framework_gemini/_sdk_client.py | 76 ++++++++++--------- .../gemini/tests/test_gemini_client.py | 21 ++++- 5 files changed, 63 insertions(+), 50 deletions(-) diff --git a/python/packages/gemini/AGENTS.md b/python/packages/gemini/AGENTS.md index b936bfea0b3..a23f3cddb5b 100644 --- a/python/packages/gemini/AGENTS.md +++ b/python/packages/gemini/AGENTS.md @@ -1,6 +1,8 @@ # Gemini Package (agent-framework-gemini) Integration with Google's Gemini Developer API and Enterprise (Vertex AI) via the `google-genai` SDK. +The shared `_sdk_client.create_genai_client` resolves authentication, backend mode, +and service URL for chat and embeddings. ## Core Classes diff --git a/python/packages/gemini/agent_framework_gemini/_chat_client.py b/python/packages/gemini/agent_framework_gemini/_chat_client.py index 18d0fc8afc5..793dc891c4f 100644 --- a/python/packages/gemini/agent_framework_gemini/_chat_client.py +++ b/python/packages/gemini/agent_framework_gemini/_chat_client.py @@ -51,8 +51,6 @@ from ._sdk_client import ( GoogleGeminiSettings, create_genai_client, - resolve_service_url, - resolve_vertexai_mode, ) if sys.version_info >= (3, 13): @@ -334,7 +332,7 @@ def __init__( configured_enterprise = google_settings.get("genai_use_enterprise") configured_vertexai = google_settings.get("genai_use_vertexai") - self._genai_client = create_genai_client( + self._genai_client, self._vertexai, self._service_url = create_genai_client( client=client, api_key=google_settings.get("api_key"), enterprise=configured_enterprise, @@ -344,9 +342,6 @@ def __init__( credentials=credentials, ) - configured_mode = configured_enterprise if configured_enterprise is not None else configured_vertexai - self._vertexai = resolve_vertexai_mode(self._genai_client, fallback=configured_mode) - self._service_url = resolve_service_url(self._genai_client, vertexai=self._vertexai) self.model = google_settings.get("model") self.max_tracked_thought_signatures = max_tracked_thought_signatures self._thought_signature_cache: OrderedDict[str, bytes] = OrderedDict() diff --git a/python/packages/gemini/agent_framework_gemini/_embedding_client.py b/python/packages/gemini/agent_framework_gemini/_embedding_client.py index 756b5b77eb3..fcbce481d87 100644 --- a/python/packages/gemini/agent_framework_gemini/_embedding_client.py +++ b/python/packages/gemini/agent_framework_gemini/_embedding_client.py @@ -34,8 +34,6 @@ from ._sdk_client import ( GoogleGeminiSettings, create_genai_client, - resolve_service_url, - resolve_vertexai_mode, ) if sys.version_info >= (3, 13): @@ -182,7 +180,7 @@ def __init__( configured_enterprise = google_settings.get("genai_use_enterprise") configured_vertexai = google_settings.get("genai_use_vertexai") - self._genai_client = create_genai_client( + self._genai_client, self._vertexai, self._service_url = create_genai_client( client=client, api_key=google_settings.get("api_key"), enterprise=configured_enterprise, @@ -192,9 +190,6 @@ def __init__( credentials=credentials, ) self._owns_client = client is None - configured_mode = configured_enterprise if configured_enterprise is not None else configured_vertexai - self._vertexai = resolve_vertexai_mode(self._genai_client, fallback=configured_mode) - self._service_url = resolve_service_url(self._genai_client, vertexai=self._vertexai) super().__init__(additional_properties=additional_properties) diff --git a/python/packages/gemini/agent_framework_gemini/_sdk_client.py b/python/packages/gemini/agent_framework_gemini/_sdk_client.py index bb7610f52b7..f4978a15463 100644 --- a/python/packages/gemini/agent_framework_gemini/_sdk_client.py +++ b/python/packages/gemini/agent_framework_gemini/_sdk_client.py @@ -88,40 +88,42 @@ def create_genai_client( location: str | None, credentials: Credentials | None, enterprise: bool | None = None, -) -> genai.Client: - """Create an SDK client for either Gemini Developer API or Vertex AI, unless one was supplied.""" - if client is not None: - return client - - if api_key is not None and not api_key.get_secret_value().strip(): - raise ValueError("GOOGLE_API_KEY must not be empty when provided.") - if enterprise is not None and vertexai is not None and enterprise != vertexai: - raise ValueError("GOOGLE_GENAI_USE_ENTERPRISE and GOOGLE_GENAI_USE_VERTEXAI cannot disagree.") - use_enterprise = enterprise if enterprise is not None else vertexai - _validate_client_auth_configuration( - vertexai=use_enterprise, - api_key=api_key, - project=project, - location=location, - credentials=credentials, - ) - client_kwargs: dict[str, Any] = { - "http_options": {"headers": {"x-goog-api-client": get_user_agent()}}, - } - if enterprise is not None: - client_kwargs["enterprise"] = enterprise - elif vertexai is not None: - client_kwargs["vertexai"] = vertexai - - if api_key is not None and (use_enterprise is not True or (credentials is None and not (project and location))): - client_kwargs["api_key"] = api_key.get_secret_value() - - if use_enterprise is True and project: - client_kwargs["project"] = project - - if use_enterprise is True and location: - client_kwargs["location"] = location - if use_enterprise is True and credentials is not None: - client_kwargs["credentials"] = credentials - - return genai.Client(**client_kwargs) +) -> tuple[genai.Client, bool, str]: + """Return the SDK client, resolved Enterprise mode, and service URL.""" + if client is None: + if api_key is not None and not api_key.get_secret_value().strip(): + raise ValueError("GOOGLE_API_KEY must not be empty when provided.") + if enterprise is not None and vertexai is not None and enterprise != vertexai: + raise ValueError("GOOGLE_GENAI_USE_ENTERPRISE and GOOGLE_GENAI_USE_VERTEXAI cannot disagree.") + use_enterprise = enterprise if enterprise is not None else vertexai + _validate_client_auth_configuration( + vertexai=use_enterprise, + api_key=api_key, + project=project, + location=location, + credentials=credentials, + ) + client_kwargs: dict[str, Any] = { + "http_options": {"headers": {"x-goog-api-client": get_user_agent()}}, + } + if enterprise is not None: + client_kwargs["enterprise"] = enterprise + elif vertexai is not None: + client_kwargs["vertexai"] = vertexai + + if api_key is not None and (use_enterprise is not True or (credentials is None and not (project and location))): + client_kwargs["api_key"] = api_key.get_secret_value() + + if use_enterprise is True and project: + client_kwargs["project"] = project + + if use_enterprise is True and location: + client_kwargs["location"] = location + if use_enterprise is True and credentials is not None: + client_kwargs["credentials"] = credentials + + client = genai.Client(**client_kwargs) + + configured_mode = enterprise if enterprise is not None else vertexai + resolved_vertexai = resolve_vertexai_mode(client, fallback=configured_mode) + return client, resolved_vertexai, resolve_service_url(client, vertexai=resolved_vertexai) diff --git a/python/packages/gemini/tests/test_gemini_client.py b/python/packages/gemini/tests/test_gemini_client.py index 53748972401..cceb1d847b8 100644 --- a/python/packages/gemini/tests/test_gemini_client.py +++ b/python/packages/gemini/tests/test_gemini_client.py @@ -23,7 +23,13 @@ from pydantic import BaseModel from typing_extensions import NotRequired, TypedDict -from agent_framework_gemini import GeminiChatClient, GeminiChatOptions, RawGeminiChatClient, ThinkingConfig +from agent_framework_gemini import ( + GeminiChatClient, + GeminiChatOptions, + GeminiEmbeddingClient, + RawGeminiChatClient, + ThinkingConfig, +) from agent_framework_gemini._feature_usage import FeatureIndex @@ -2757,6 +2763,19 @@ def test_service_url_falls_back_when_sdk_base_url_is_unavailable() -> None: assert vertex_client.service_url() == "https://aiplatform.googleapis.com" +def test_injected_sdk_routing_is_shared_with_embeddings() -> None: + sdk = MagicMock() + sdk._api_client.vertexai = True + sdk._api_client._http_options.base_url = "https://custom.example.test/" + + chat = GeminiChatClient(client=sdk, model="gemini-2.5-flash", vertexai=False) + embeddings = GeminiEmbeddingClient(client=sdk, model="gemini-embedding-2", vertexai=False) + + assert chat._genai_client is embeddings._genai_client is sdk + assert chat._vertexai is embeddings._vertexai is True + assert chat.service_url() == embeddings.service_url() == "https://custom.example.test" + + # integration tests