Skip to content

Commit 373c580

Browse files
DanielTobi0marandaneto
authored andcommitted
feat(gemini): add aio surface and files API (Fixes #315)
The Gemini adapter advertised itself as a drop-in replacement for genai.Client, but two surfaces of the real SDK were missing: - `client.aio.models.*` - async generation was only reachable by swapping the class out for `AsyncClient`, so SDK-shaped code hit an AttributeError on `aio`. - `client.files.*` - absent entirely, which breaks any multimodal flow that uploads a file before referencing it in `contents`. `Client` now builds the provider client once and shares it across every surface, so `models`, `aio.models` and `files` no longer open separate connections. `aio.models` is a tracked `AsyncModels` inheriting the same PostHog defaults; `files` and `aio.files` pass straight through to the provider, since uploads are not generations and emit no events. `AsyncClient` gains the matching pair: `files` resolves to the provider's async Files API, and `aio` aliases its already-async `models`.
1 parent f1b6e29 commit 373c580

6 files changed

Lines changed: 251 additions & 18 deletions

File tree

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
---
2+
pypi/posthog: minor
3+
---
4+
5+
The Gemini adapter now covers two surfaces of `genai.Client` it previously lacked. `Client.aio.models` reaches the tracked async models adapter, so `await client.aio.models.generate_content(...)` works without swapping the class out for `AsyncClient`, and `client.files` (plus `client.aio.files`, and `AsyncClient.files` for the async Files API) passes through to the provider, so multimodal flows that upload a file before referencing it in `contents` no longer fail. Every surface of one client shares a single provider client instead of opening its own.

‎posthog/ai/gemini/_shared.py‎

Lines changed: 53 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -78,6 +78,44 @@ def _build_gemini_client_args(
7878
return client_args
7979

8080

81+
def _build_gemini_client(
82+
*,
83+
api_key: Optional[str],
84+
vertexai: Optional[bool],
85+
credentials: Optional[Any],
86+
project: Optional[str],
87+
location: Optional[str],
88+
debug_config: Optional[Any],
89+
http_options: Optional[Any],
90+
) -> Any:
91+
"""Construct the provider client shared by every surface of one adapter."""
92+
return genai.Client(
93+
**_build_gemini_client_args(
94+
api_key=api_key,
95+
vertexai=vertexai,
96+
credentials=credentials,
97+
project=project,
98+
location=location,
99+
debug_config=debug_config,
100+
http_options=http_options,
101+
)
102+
)
103+
104+
105+
class _GeminiAioNamespace:
106+
"""
107+
Mirrors ``genai.Client().aio``: the async surface of a single client.
108+
109+
Pairs the tracked async ``models`` adapter with the provider's async
110+
``files`` API so ``client.aio.models`` and ``client.aio.files`` resolve the
111+
same way they do on the real SDK.
112+
"""
113+
114+
def __init__(self, models: Any, files: Any):
115+
self.models = models
116+
self.files = files
117+
118+
81119
class _GeminiModelsPolicy:
82120
"""Shared telemetry policy for the explicit sync and async Gemini adapters."""
83121

@@ -98,23 +136,29 @@ def _initialize_policy(
98136
posthog_properties: Optional[Dict[str, Any]],
99137
posthog_privacy_mode: bool,
100138
posthog_groups: Optional[Dict[str, Any]],
139+
provider_client: Optional[Any] = None,
101140
) -> None:
102141
self._ph_client = _resolve_posthog_client(posthog_client)
103142
self._default_distinct_id = posthog_distinct_id
104143
self._default_properties = posthog_properties or {}
105144
self._default_privacy_mode = posthog_privacy_mode
106145
self._default_groups = posthog_groups
107146

108-
client_args = _build_gemini_client_args(
109-
api_key=api_key,
110-
vertexai=vertexai,
111-
credentials=credentials,
112-
project=project,
113-
location=location,
114-
debug_config=debug_config,
115-
http_options=http_options,
147+
# Every surface of one PostHog client (sync models, aio models, files)
148+
# shares a single provider client rather than opening its own.
149+
self._client = (
150+
provider_client
151+
if provider_client is not None
152+
else _build_gemini_client(
153+
api_key=api_key,
154+
vertexai=vertexai,
155+
credentials=credentials,
156+
project=project,
157+
location=location,
158+
debug_config=debug_config,
159+
http_options=http_options,
160+
)
116161
)
117-
self._client = genai.Client(**client_args)
118162
self._base_url = _GEMINI_BASE_URL
119163

120164
def _merge_posthog_params(

‎posthog/ai/gemini/gemini.py‎

Lines changed: 41 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,13 @@
1313
merge_usage_stats,
1414
with_privacy_mode as with_privacy_mode,
1515
)
16-
from ._shared import _GeminiModelsPolicy, _resolve_posthog_client
16+
from ._shared import (
17+
_GeminiAioNamespace,
18+
_GeminiModelsPolicy,
19+
_build_gemini_client,
20+
_resolve_posthog_client,
21+
)
22+
from .gemini_async import AsyncModels
1723
from .gemini_converter import (
1824
extract_gemini_content_from_chunk,
1925
extract_gemini_embedding_token_count as extract_gemini_embedding_token_count,
@@ -39,6 +45,16 @@ class Client:
3945
contents=["Hello world"],
4046
posthog_distinct_id="specific_user" # Override default
4147
)
48+
49+
The async surface lives under ``aio``, exactly as it does on genai.Client:
50+
51+
response = await client.aio.models.generate_content(
52+
model="gemini-2.0-flash",
53+
contents=["Hello world"],
54+
)
55+
56+
``files`` (and ``aio.files``) pass straight through to the provider. Uploads
57+
are not generations, so they emit no PostHog events.
4258
"""
4359

4460
_ph_client: PostHogClient
@@ -78,21 +94,40 @@ def __init__(
7894

7995
self._ph_client = _resolve_posthog_client(posthog_client)
8096

81-
self.models = Models(
97+
# Built once here and shared by every surface, so a client that uses both
98+
# `models` and `aio.models` still opens a single provider client.
99+
self._provider_client = _build_gemini_client(
82100
api_key=api_key,
83101
vertexai=vertexai,
84102
credentials=credentials,
85103
project=project,
86104
location=location,
87105
debug_config=debug_config,
88106
http_options=http_options,
107+
)
108+
109+
self.models = Models(
110+
provider_client=self._provider_client,
89111
posthog_client=self._ph_client,
90112
posthog_distinct_id=posthog_distinct_id,
91113
posthog_properties=posthog_properties,
92114
posthog_privacy_mode=posthog_privacy_mode,
93115
posthog_groups=posthog_groups,
94116
**kwargs,
95117
)
118+
self.files = self._provider_client.files
119+
self.aio = _GeminiAioNamespace(
120+
models=AsyncModels(
121+
provider_client=self._provider_client,
122+
posthog_client=self._ph_client,
123+
posthog_distinct_id=posthog_distinct_id,
124+
posthog_properties=posthog_properties,
125+
posthog_privacy_mode=posthog_privacy_mode,
126+
posthog_groups=posthog_groups,
127+
**kwargs,
128+
),
129+
files=self._provider_client.aio.files,
130+
)
96131

97132

98133
class Models(_GeminiModelsPolicy):
@@ -116,6 +151,7 @@ def __init__(
116151
posthog_properties: Optional[Dict[str, Any]] = None,
117152
posthog_privacy_mode: bool = False,
118153
posthog_groups: Optional[Dict[str, Any]] = None,
154+
provider_client: Optional[Any] = None,
119155
**kwargs,
120156
):
121157
"""
@@ -132,6 +168,8 @@ def __init__(
132168
posthog_properties: Default properties for all calls
133169
posthog_privacy_mode: Default privacy mode for all calls
134170
posthog_groups: Default groups for all calls
171+
provider_client: An already-built genai.Client to reuse instead of
172+
constructing one from the connection arguments above
135173
**kwargs: Additional arguments (for future compatibility)
136174
"""
137175

@@ -148,6 +186,7 @@ def __init__(
148186
posthog_properties=posthog_properties,
149187
posthog_privacy_mode=posthog_privacy_mode,
150188
posthog_groups=posthog_groups,
189+
provider_client=provider_client,
151190
)
152191

153192
def generate_content(

‎posthog/ai/gemini/gemini_async.py‎

Lines changed: 24 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,12 @@
1414
merge_usage_stats,
1515
with_privacy_mode as with_privacy_mode,
1616
)
17-
from ._shared import _GeminiModelsPolicy, _resolve_posthog_client
17+
from ._shared import (
18+
_GeminiAioNamespace,
19+
_GeminiModelsPolicy,
20+
_build_gemini_client,
21+
_resolve_posthog_client,
22+
)
1823
from .gemini_converter import (
1924
extract_gemini_content_from_chunk,
2025
extract_gemini_embedding_token_count as extract_gemini_embedding_token_count,
@@ -40,6 +45,11 @@ class AsyncClient:
4045
contents=["Hello world"],
4146
posthog_distinct_id="specific_user" # Override default
4247
)
48+
49+
``models`` is already the async surface, and ``aio.models`` is an alias for
50+
it so code copied from the Google SDK docs keeps working. ``files`` (and
51+
``aio.files``) pass straight through to the provider's async Files API;
52+
uploads are not generations, so they emit no PostHog events.
4353
"""
4454

4555
_ph_client: PostHogClient
@@ -79,21 +89,29 @@ def __init__(
7989

8090
self._ph_client = _resolve_posthog_client(posthog_client)
8191

82-
self.models = AsyncModels(
92+
# Built once here and shared by every surface, so `models`, `aio.models`
93+
# and `files` all go through a single provider client.
94+
self._provider_client = _build_gemini_client(
8395
api_key=api_key,
8496
vertexai=vertexai,
8597
credentials=credentials,
8698
project=project,
8799
location=location,
88100
debug_config=debug_config,
89101
http_options=http_options,
102+
)
103+
104+
self.models = AsyncModels(
105+
provider_client=self._provider_client,
90106
posthog_client=self._ph_client,
91107
posthog_distinct_id=posthog_distinct_id,
92108
posthog_properties=posthog_properties,
93109
posthog_privacy_mode=posthog_privacy_mode,
94110
posthog_groups=posthog_groups,
95111
**kwargs,
96112
)
113+
self.files = self._provider_client.aio.files
114+
self.aio = _GeminiAioNamespace(models=self.models, files=self.files)
97115

98116

99117
class AsyncModels(_GeminiModelsPolicy):
@@ -117,6 +135,7 @@ def __init__(
117135
posthog_properties: Optional[Dict[str, Any]] = None,
118136
posthog_privacy_mode: bool = False,
119137
posthog_groups: Optional[Dict[str, Any]] = None,
138+
provider_client: Optional[Any] = None,
120139
**kwargs,
121140
):
122141
"""
@@ -133,6 +152,8 @@ def __init__(
133152
posthog_properties: Default properties for all calls
134153
posthog_privacy_mode: Default privacy mode for all calls
135154
posthog_groups: Default groups for all calls
155+
provider_client: An already-built genai.Client to reuse instead of
156+
constructing one from the connection arguments above
136157
**kwargs: Additional arguments (for future compatibility)
137158
"""
138159

@@ -149,6 +170,7 @@ def __init__(
149170
posthog_properties=posthog_properties,
150171
posthog_privacy_mode=posthog_privacy_mode,
151172
posthog_groups=posthog_groups,
173+
provider_client=provider_client,
152174
)
153175

154176
async def generate_content(

‎posthog/test/ai/gemini/test_gemini_parity.py‎

Lines changed: 119 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
from unittest.mock import MagicMock, patch
1+
from unittest.mock import AsyncMock, MagicMock, patch
22

33
import pytest
44

@@ -65,3 +65,121 @@ def test_sync_and_async_clients_merge_posthog_defaults_without_mutation(client_c
6565
{"organization": "default-org"},
6666
)
6767
assert default_properties == {"shared": "default", "default-only": True}
68+
69+
70+
@pytest.fixture
71+
def provider_client_class():
72+
with patch.object(google_genai, "Client") as patched:
73+
patched.return_value = MagicMock()
74+
yield patched
75+
76+
77+
@pytest.fixture
78+
def provider_client(provider_client_class):
79+
"""A stand-in for the underlying google-genai client, with both surfaces."""
80+
return provider_client_class.return_value
81+
82+
83+
@pytest.fixture
84+
def gemini_response():
85+
response = MagicMock()
86+
response.text = "Test response from Gemini"
87+
88+
usage = MagicMock()
89+
usage.prompt_token_count = 20
90+
usage.candidates_token_count = 10
91+
usage.cached_content_token_count = 0
92+
usage.thoughts_token_count = 0
93+
response.usage_metadata = usage
94+
95+
part = MagicMock()
96+
part.text = "Test response from Gemini"
97+
content = MagicMock()
98+
content.parts = [part]
99+
candidate = MagicMock()
100+
candidate.content = content
101+
response.candidates = [candidate]
102+
103+
return response
104+
105+
106+
@pytest.mark.parametrize("client_class", [Client, AsyncClient])
107+
def test_every_surface_shares_one_provider_client(
108+
client_class, provider_client_class, provider_client
109+
):
110+
"""`models`, `aio.models` and `files` must not open separate connections."""
111+
client = client_class(api_key="test-key", posthog_client=MagicMock())
112+
113+
provider_client_class.assert_called_once()
114+
assert client.models._client is provider_client
115+
assert client.aio.models._client is provider_client
116+
117+
118+
def test_sync_client_exposes_the_provider_files_api(provider_client):
119+
client = Client(api_key="test-key", posthog_client=MagicMock())
120+
121+
assert client.files is provider_client.files
122+
assert client.aio.files is provider_client.aio.files
123+
124+
uploaded = client.files.upload(file="notes.pdf")
125+
126+
provider_client.files.upload.assert_called_once_with(file="notes.pdf")
127+
assert uploaded is provider_client.files.upload.return_value
128+
129+
130+
def test_async_client_exposes_the_async_files_api(provider_client):
131+
"""AsyncClient is async end to end, so its files API is the aio one."""
132+
client = AsyncClient(api_key="test-key", posthog_client=MagicMock())
133+
134+
assert client.files is provider_client.aio.files
135+
assert client.aio.files is provider_client.aio.files
136+
# `models` is already async; `aio.models` is an alias so SDK-shaped code works.
137+
assert client.aio.models is client.models
138+
139+
140+
def test_sync_client_aio_models_inherits_posthog_defaults(provider_client):
141+
client = Client(
142+
api_key="test-key",
143+
posthog_client=MagicMock(),
144+
posthog_distinct_id="default-id",
145+
posthog_properties={"team": "ai"},
146+
posthog_privacy_mode=True,
147+
posthog_groups={"organization": "default-org"},
148+
)
149+
150+
assert client.aio.models._merge_posthog_params(None, "trace", None, None, None) == (
151+
"default-id",
152+
"trace",
153+
{"team": "ai"},
154+
True,
155+
{"organization": "default-org"},
156+
)
157+
158+
159+
@pytest.mark.asyncio
160+
async def test_sync_client_aio_models_tracks_generations(
161+
provider_client, gemini_response
162+
):
163+
"""`client.aio.models.generate_content` is the drop-in async entry point."""
164+
provider_client.aio.models.generate_content = AsyncMock(
165+
return_value=gemini_response
166+
)
167+
posthog_client = MagicMock()
168+
posthog_client.privacy_mode = False
169+
170+
client = Client(api_key="test-key", posthog_client=posthog_client)
171+
172+
response = await client.aio.models.generate_content(
173+
model="gemini-2.0-flash",
174+
contents=["Tell me a fun fact about hedgehogs"],
175+
posthog_distinct_id="test-id",
176+
)
177+
178+
assert response is gemini_response
179+
provider_client.aio.models.generate_content.assert_awaited_once()
180+
181+
assert posthog_client.capture.call_count == 1
182+
call_args = posthog_client.capture.call_args[1]
183+
assert call_args["distinct_id"] == "test-id"
184+
assert call_args["event"] == "$ai_generation"
185+
assert call_args["properties"]["$ai_model"] == "gemini-2.0-flash"

0 commit comments

Comments
 (0)