Compare commits

...
4 Commits
Author SHA1 Message Date
Radu RaiceaandGitHub 68e78c877d feat(llmo): support Vertex AI (#302)
* feat(llmo): support Vertex AI

* chore(llmo): run formatter

* fix(llmo): fix types error

* chore(llmo): run formatter

* chore(llmo): bump version
2025-08-05 15:33:10 -04:00
Radu RaiceaandGitHub 07cf32bb04 fix(llmo): tool calls are broken for most providers (#299)
* fix(llmo): set the $ai_tools properly for all providers

* fix(llmo): remove privacy mode from $ai_tools

* chore(llmo): bump version

* chore(llmo): run formatter

* fix(llmo): properly set tool calls in $ai_output_choices

* chore(llmo): bump version

* chore(llmo): run formatter

* fix(llmo): fix types error

* feat(llmo): change $ai_output_choices to have an array of content

* chore(llmo): run formatter

* feat(llmo): create text type object

* chore(llmo): update CHANGELOG.md
2025-08-05 14:02:37 -04:00
Phil HaackGitHubgreptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com>
0076b66b75 feat: Expose get_feature_flag_result method in public API (#284)
Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com>
2025-08-05 10:14:29 -07:00
Dylan MartinandGitHub 09dad8117f fix (#300) 2025-08-01 17:32:58 -07:00
16 changed files with 1497 additions and 317 deletions
+12
View File
@@ -1,3 +1,15 @@
# 6.4.0 - 2025-08-05
- feat: support Vertex AI for Gemini
# 6.3.4 - 2025-08-04
- fix: set `$ai_tools` for all providers and `$ai_output_choices` for all non-streaming provider flows properly
# 6.3.3 - 2025-08-01
- fix: `get_feature_flag_result` now correctly returns FeatureFlagResult when payload is empty string instead of None
# 6.3.2 - 2025-07-31
- fix: Anthropic's tool calls are now handled properly
+13 -2
View File
@@ -59,8 +59,19 @@ print(posthog.feature_enabled("beta-feature", "distinct_id"))
# get payload
print(posthog.get_feature_flag_payload("beta-feature", "distinct_id"))
print(posthog.get_all_flags_and_payloads("distinct_id"))
exit()
# # Alias a previous distinct id with a new one
# get feature flag result with all details (enabled, variant, payload, key, reason)
result = posthog.get_feature_flag_result("beta-feature", "distinct_id")
if result:
print(f"Flag key: {result.key}")
print(f"Flag enabled: {result.enabled}")
print(f"Variant: {result.variant}")
print(f"Payload: {result.payload}")
print(f"Reason: {result.reason}")
# get_value() returns the variant if it exists, otherwise the enabled value
print(f"Value (variant or enabled): {result.get_value()}")
# Alias a previous distinct id with a new one
posthog.alias("distinct_id", "new_distinct_id")
+74 -31
View File
@@ -11,7 +11,7 @@ from posthog.contexts import (
set_context_session as inner_set_context_session,
identify_context as inner_identify_context,
)
from posthog.types import FeatureFlag, FlagsAndPayloads
from posthog.types import FeatureFlag, FlagsAndPayloads, FeatureFlagResult
from posthog.version import VERSION
__version__ = VERSION
@@ -388,9 +388,9 @@ def capture_exception(
def feature_enabled(
key, # type: str
distinct_id, # type: str
groups={}, # type: dict
person_properties={}, # type: dict
group_properties={}, # type: dict
groups=None, # type: Optional[dict]
person_properties=None, # type: Optional[dict]
group_properties=None, # type: Optional[dict]
only_evaluate_locally=False, # type: bool
send_feature_flag_events=True, # type: bool
disable_geoip=None, # type: Optional[bool]
@@ -427,9 +427,9 @@ def feature_enabled(
"feature_enabled",
key=key,
distinct_id=distinct_id,
groups=groups,
person_properties=person_properties,
group_properties=group_properties,
groups=groups or {},
person_properties=person_properties or {},
group_properties=group_properties or {},
only_evaluate_locally=only_evaluate_locally,
send_feature_flag_events=send_feature_flag_events,
disable_geoip=disable_geoip,
@@ -439,9 +439,9 @@ def feature_enabled(
def get_feature_flag(
key, # type: str
distinct_id, # type: str
groups={}, # type: dict
person_properties={}, # type: dict
group_properties={}, # type: dict
groups=None, # type: Optional[dict]
person_properties=None, # type: Optional[dict]
group_properties=None, # type: Optional[dict]
only_evaluate_locally=False, # type: bool
send_feature_flag_events=True, # type: bool
disable_geoip=None, # type: Optional[bool]
@@ -477,9 +477,9 @@ def get_feature_flag(
"get_feature_flag",
key=key,
distinct_id=distinct_id,
groups=groups,
person_properties=person_properties,
group_properties=group_properties,
groups=groups or {},
person_properties=person_properties or {},
group_properties=group_properties or {},
only_evaluate_locally=only_evaluate_locally,
send_feature_flag_events=send_feature_flag_events,
disable_geoip=disable_geoip,
@@ -488,9 +488,9 @@ def get_feature_flag(
def get_all_flags(
distinct_id, # type: str
groups={}, # type: dict
person_properties={}, # type: dict
group_properties={}, # type: dict
groups=None, # type: Optional[dict]
person_properties=None, # type: Optional[dict]
group_properties=None, # type: Optional[dict]
only_evaluate_locally=False, # type: bool
disable_geoip=None, # type: Optional[bool]
) -> Optional[dict[str, FeatureFlag]]:
@@ -520,21 +520,64 @@ def get_all_flags(
return _proxy(
"get_all_flags",
distinct_id=distinct_id,
groups=groups,
person_properties=person_properties,
group_properties=group_properties,
groups=groups or {},
person_properties=person_properties or {},
group_properties=group_properties or {},
only_evaluate_locally=only_evaluate_locally,
disable_geoip=disable_geoip,
)
def get_feature_flag_result(
key,
distinct_id,
groups=None, # type: Optional[dict]
person_properties=None, # type: Optional[dict]
group_properties=None, # type: Optional[dict]
only_evaluate_locally=False,
send_feature_flag_events=True,
disable_geoip=None, # type: Optional[bool]
):
# type: (...) -> Optional[FeatureFlagResult]
"""
Get a FeatureFlagResult object which contains the flag result and payload.
This method evaluates a feature flag and returns a FeatureFlagResult object containing:
- enabled: Whether the flag is enabled
- variant: The variant value if the flag has variants
- payload: The payload associated with the flag (automatically deserialized from JSON)
- key: The flag key
- reason: Why the flag was enabled/disabled
Example:
```python
result = posthog.get_feature_flag_result('beta-feature', 'distinct_id')
if result and result.enabled:
# Use the variant and payload
print(f"Variant: {result.variant}")
print(f"Payload: {result.payload}")
```
"""
return _proxy(
"get_feature_flag_result",
key=key,
distinct_id=distinct_id,
groups=groups or {},
person_properties=person_properties or {},
group_properties=group_properties or {},
only_evaluate_locally=only_evaluate_locally,
send_feature_flag_events=send_feature_flag_events,
disable_geoip=disable_geoip,
)
def get_feature_flag_payload(
key,
distinct_id,
match_value=None,
groups={},
person_properties={},
group_properties={},
groups=None, # type: Optional[dict]
person_properties=None, # type: Optional[dict]
group_properties=None, # type: Optional[dict]
only_evaluate_locally=False,
send_feature_flag_events=True,
disable_geoip=None, # type: Optional[bool]
@@ -544,9 +587,9 @@ def get_feature_flag_payload(
key=key,
distinct_id=distinct_id,
match_value=match_value,
groups=groups,
person_properties=person_properties,
group_properties=group_properties,
groups=groups or {},
person_properties=person_properties or {},
group_properties=group_properties or {},
only_evaluate_locally=only_evaluate_locally,
send_feature_flag_events=send_feature_flag_events,
disable_geoip=disable_geoip,
@@ -575,18 +618,18 @@ def get_remote_config_payload(
def get_all_flags_and_payloads(
distinct_id,
groups={},
person_properties={},
group_properties={},
groups=None, # type: Optional[dict]
person_properties=None, # type: Optional[dict]
group_properties=None, # type: Optional[dict]
only_evaluate_locally=False,
disable_geoip=None, # type: Optional[bool]
) -> FlagsAndPayloads:
return _proxy(
"get_all_flags_and_payloads",
distinct_id=distinct_id,
groups=groups,
person_properties=person_properties,
group_properties=group_properties,
groups=groups or {},
person_properties=person_properties or {},
group_properties=group_properties or {},
only_evaluate_locally=only_evaluate_locally,
disable_geoip=disable_geoip,
)
+64 -10
View File
@@ -42,6 +42,12 @@ class Client:
def __init__(
self,
api_key: Optional[str] = None,
vertexai: Optional[bool] = None,
credentials: Optional[Any] = None,
project: Optional[str] = None,
location: Optional[str] = None,
debug_config: Optional[Any] = None,
http_options: Optional[Any] = None,
posthog_client: Optional[PostHogClient] = None,
posthog_distinct_id: Optional[str] = None,
posthog_properties: Optional[Dict[str, Any]] = None,
@@ -51,7 +57,13 @@ class Client:
):
"""
Args:
api_key: Google AI API key. If not provided, will use GOOGLE_API_KEY or API_KEY environment variable
api_key: Google AI API key. If not provided, will use GOOGLE_API_KEY or API_KEY environment variable (not required for Vertex AI)
vertexai: Whether to use Vertex AI authentication
credentials: Vertex AI credentials object
project: GCP project ID for Vertex AI
location: GCP location for Vertex AI
debug_config: Debug configuration for the client
http_options: HTTP options for the client
posthog_client: PostHog client for tracking usage
posthog_distinct_id: Default distinct ID for all calls (can be overridden per call)
posthog_properties: Default properties for all calls (can be overridden per call)
@@ -66,6 +78,12 @@ class Client:
self.models = Models(
api_key=api_key,
vertexai=vertexai,
credentials=credentials,
project=project,
location=location,
debug_config=debug_config,
http_options=http_options,
posthog_client=self._ph_client,
posthog_distinct_id=posthog_distinct_id,
posthog_properties=posthog_properties,
@@ -85,6 +103,12 @@ class Models:
def __init__(
self,
api_key: Optional[str] = None,
vertexai: Optional[bool] = None,
credentials: Optional[Any] = None,
project: Optional[str] = None,
location: Optional[str] = None,
debug_config: Optional[Any] = None,
http_options: Optional[Any] = None,
posthog_client: Optional[PostHogClient] = None,
posthog_distinct_id: Optional[str] = None,
posthog_properties: Optional[Dict[str, Any]] = None,
@@ -94,7 +118,13 @@ class Models:
):
"""
Args:
api_key: Google AI API key. If not provided, will use GOOGLE_API_KEY or API_KEY environment variable
api_key: Google AI API key. If not provided, will use GOOGLE_API_KEY or API_KEY environment variable (not required for Vertex AI)
vertexai: Whether to use Vertex AI authentication
credentials: Vertex AI credentials object
project: GCP project ID for Vertex AI
location: GCP location for Vertex AI
debug_config: Debug configuration for the client
http_options: HTTP options for the client
posthog_client: PostHog client for tracking usage
posthog_distinct_id: Default distinct ID for all calls
posthog_properties: Default properties for all calls
@@ -113,16 +143,40 @@ class Models:
self._default_privacy_mode = posthog_privacy_mode
self._default_groups = posthog_groups
# Handle API key - try parameter first, then environment variables
if api_key is None:
api_key = os.environ.get("GOOGLE_API_KEY") or os.environ.get("API_KEY")
# Build genai.Client arguments
client_args: Dict[str, Any] = {}
if api_key is None:
raise ValueError(
"API key must be provided either as parameter or via GOOGLE_API_KEY/API_KEY environment variable"
)
# Add Vertex AI parameters if provided
if vertexai is not None:
client_args["vertexai"] = vertexai
if credentials is not None:
client_args["credentials"] = credentials
if project is not None:
client_args["project"] = project
if location is not None:
client_args["location"] = location
if debug_config is not None:
client_args["debug_config"] = debug_config
if http_options is not None:
client_args["http_options"] = http_options
self._client = genai.Client(api_key=api_key)
# Handle API key authentication
if vertexai:
# For Vertex AI, api_key is optional
if api_key is not None:
client_args["api_key"] = api_key
else:
# For non-Vertex AI mode, api_key is required (backwards compatibility)
if api_key is None:
api_key = os.environ.get("GOOGLE_API_KEY") or os.environ.get("API_KEY")
if api_key is None:
raise ValueError(
"API key must be provided either as parameter or via GOOGLE_API_KEY/API_KEY environment variable"
)
client_args["api_key"] = api_key
self._client = genai.Client(**client_args)
self._base_url = "https://generativelanguage.googleapis.com"
def _merge_posthog_params(
+5 -7
View File
@@ -556,12 +556,9 @@ class CallbackHandler(BaseCallbackHandler):
"$ai_latency": run.latency,
"$ai_base_url": run.base_url,
}
if run.tools:
event_properties["$ai_tools"] = with_privacy_mode(
self._ph_client,
self._privacy_mode,
run.tools,
)
event_properties["$ai_tools"] = run.tools
if isinstance(output, BaseException):
event_properties["$ai_http_status"] = _get_http_status(output)
@@ -587,7 +584,8 @@ class CallbackHandler(BaseCallbackHandler):
]
else:
completions = [
_extract_raw_esponse(generation) for generation in generation_result
_extract_raw_response(generation)
for generation in generation_result
]
event_properties["$ai_output_choices"] = with_privacy_mode(
self._ph_client, self._privacy_mode, completions
@@ -618,7 +616,7 @@ class CallbackHandler(BaseCallbackHandler):
)
def _extract_raw_esponse(last_response):
def _extract_raw_response(last_response):
"""Extract the response from the last response of the LLM call."""
# We return the text of the response if not empty
if last_response.text is not None and last_response.text.strip() != "":
+9 -36
View File
@@ -11,6 +11,7 @@ except ImportError:
from posthog.ai.utils import (
call_llm_and_track_usage,
extract_available_tool_calls,
get_model_params,
with_privacy_mode,
)
@@ -167,6 +168,7 @@ class WrappedResponses:
usage_stats,
latency,
output,
extract_available_tool_calls("openai", kwargs),
)
return generator()
@@ -182,7 +184,7 @@ class WrappedResponses:
usage_stats: Dict[str, int],
latency: float,
output: Any,
tool_calls: Optional[List[Dict[str, Any]]] = None,
available_tool_calls: Optional[List[Dict[str, Any]]] = None,
):
if posthog_trace_id is None:
posthog_trace_id = str(uuid.uuid4())
@@ -212,12 +214,8 @@ class WrappedResponses:
**(posthog_properties or {}),
}
if tool_calls:
event_properties["$ai_tools"] = with_privacy_mode(
self._client._ph_client,
posthog_privacy_mode,
tool_calls,
)
if available_tool_calls:
event_properties["$ai_tools"] = available_tool_calls
if posthog_distinct_id is None:
event_properties["$process_person_profile"] = False
@@ -341,7 +339,6 @@ class WrappedCompletions:
start_time = time.time()
usage_stats: Dict[str, int] = {}
accumulated_content = []
accumulated_tools = {}
if "stream_options" not in kwargs:
kwargs["stream_options"] = {}
kwargs["stream_options"]["include_usage"] = True
@@ -350,7 +347,6 @@ class WrappedCompletions:
def generator():
nonlocal usage_stats
nonlocal accumulated_content # noqa: F824
nonlocal accumulated_tools # noqa: F824
try:
for chunk in response:
@@ -389,31 +385,12 @@ class WrappedCompletions:
if content:
accumulated_content.append(content)
# Process tool calls
tool_calls = getattr(chunk.choices[0].delta, "tool_calls", None)
if tool_calls:
for tool_call in tool_calls:
index = tool_call.index
if index not in accumulated_tools:
accumulated_tools[index] = tool_call
else:
# Append arguments for existing tool calls
if hasattr(tool_call, "function") and hasattr(
tool_call.function, "arguments"
):
accumulated_tools[
index
].function.arguments += (
tool_call.function.arguments
)
yield chunk
finally:
end_time = time.time()
latency = end_time - start_time
output = "".join(accumulated_content)
tools = list(accumulated_tools.values()) if accumulated_tools else None
self._capture_streaming_event(
posthog_distinct_id,
posthog_trace_id,
@@ -424,7 +401,7 @@ class WrappedCompletions:
usage_stats,
latency,
output,
tools,
extract_available_tool_calls("openai", kwargs),
)
return generator()
@@ -440,7 +417,7 @@ class WrappedCompletions:
usage_stats: Dict[str, int],
latency: float,
output: Any,
tool_calls: Optional[List[Dict[str, Any]]] = None,
available_tool_calls: Optional[List[Dict[str, Any]]] = None,
):
if posthog_trace_id is None:
posthog_trace_id = str(uuid.uuid4())
@@ -470,12 +447,8 @@ class WrappedCompletions:
**(posthog_properties or {}),
}
if tool_calls:
event_properties["$ai_tools"] = with_privacy_mode(
self._client._ph_client,
posthog_privacy_mode,
tool_calls,
)
if available_tool_calls:
event_properties["$ai_tools"] = available_tool_calls
if posthog_distinct_id is None:
event_properties["$process_person_profile"] = False
+9 -36
View File
@@ -12,6 +12,7 @@ except ImportError:
from posthog import setup
from posthog.ai.utils import (
call_llm_and_track_usage_async,
extract_available_tool_calls,
get_model_params,
with_privacy_mode,
)
@@ -168,6 +169,7 @@ class WrappedResponses:
usage_stats,
latency,
output,
extract_available_tool_calls("openai", kwargs),
)
return async_generator()
@@ -183,7 +185,7 @@ class WrappedResponses:
usage_stats: Dict[str, int],
latency: float,
output: Any,
tool_calls: Optional[List[Dict[str, Any]]] = None,
available_tool_calls: Optional[List[Dict[str, Any]]] = None,
):
if posthog_trace_id is None:
posthog_trace_id = str(uuid.uuid4())
@@ -213,12 +215,8 @@ class WrappedResponses:
**(posthog_properties or {}),
}
if tool_calls:
event_properties["$ai_tools"] = with_privacy_mode(
self._client._ph_client,
posthog_privacy_mode,
tool_calls,
)
if available_tool_calls:
event_properties["$ai_tools"] = available_tool_calls
if posthog_distinct_id is None:
event_properties["$process_person_profile"] = False
@@ -344,7 +342,6 @@ class WrappedCompletions:
start_time = time.time()
usage_stats: Dict[str, int] = {}
accumulated_content = []
accumulated_tools = {}
if "stream_options" not in kwargs:
kwargs["stream_options"] = {}
@@ -354,7 +351,6 @@ class WrappedCompletions:
async def async_generator():
nonlocal usage_stats
nonlocal accumulated_content # noqa: F824
nonlocal accumulated_tools # noqa: F824
try:
async for chunk in response:
@@ -393,31 +389,12 @@ class WrappedCompletions:
if content:
accumulated_content.append(content)
# Process tool calls
tool_calls = getattr(chunk.choices[0].delta, "tool_calls", None)
if tool_calls:
for tool_call in tool_calls:
index = tool_call.index
if index not in accumulated_tools:
accumulated_tools[index] = tool_call
else:
# Append arguments for existing tool calls
if hasattr(tool_call, "function") and hasattr(
tool_call.function, "arguments"
):
accumulated_tools[
index
].function.arguments += (
tool_call.function.arguments
)
yield chunk
finally:
end_time = time.time()
latency = end_time - start_time
output = "".join(accumulated_content)
tools = list(accumulated_tools.values()) if accumulated_tools else None
await self._capture_streaming_event(
posthog_distinct_id,
posthog_trace_id,
@@ -428,7 +405,7 @@ class WrappedCompletions:
usage_stats,
latency,
output,
tools,
extract_available_tool_calls("openai", kwargs),
)
return async_generator()
@@ -444,7 +421,7 @@ class WrappedCompletions:
usage_stats: Dict[str, int],
latency: float,
output: Any,
tool_calls: Optional[List[Dict[str, Any]]] = None,
available_tool_calls: Optional[List[Dict[str, Any]]] = None,
):
if posthog_trace_id is None:
posthog_trace_id = str(uuid.uuid4())
@@ -474,12 +451,8 @@ class WrappedCompletions:
**(posthog_properties or {}),
}
if tool_calls:
event_properties["$ai_tools"] = with_privacy_mode(
self._client._ph_client,
posthog_privacy_mode,
tool_calls,
)
if available_tool_calls:
event_properties["$ai_tools"] = available_tool_calls
if posthog_distinct_id is None:
event_properties["$process_person_profile"] = False
+133 -90
View File
@@ -117,6 +117,8 @@ def format_response(response, provider: str):
def format_response_anthropic(response):
output = []
content = []
for choice in response.content:
if (
hasattr(choice, "type")
@@ -124,32 +126,78 @@ def format_response_anthropic(response):
and hasattr(choice, "text")
and choice.text
):
output.append(
{
"role": "assistant",
"content": choice.text,
}
)
content.append({"type": "text", "text": choice.text})
elif (
hasattr(choice, "type")
and choice.type == "tool_use"
and hasattr(choice, "name")
and hasattr(choice, "id")
):
tool_call = {
"type": "function",
"id": choice.id,
"function": {
"name": choice.name,
"arguments": getattr(choice, "input", {}),
},
}
content.append(tool_call)
if content:
message = {
"role": "assistant",
"content": content,
}
output.append(message)
return output
def format_response_openai(response):
output = []
if hasattr(response, "choices"):
content = []
role = "assistant"
for choice in response.choices:
# Handle Chat Completions response format
if hasattr(choice, "message") and choice.message and choice.message.content:
output.append(
{
"content": choice.message.content,
"role": choice.message.role,
}
)
if hasattr(choice, "message") and choice.message:
if choice.message.role:
role = choice.message.role
if choice.message.content:
content.append({"type": "text", "text": choice.message.content})
if hasattr(choice.message, "tool_calls") and choice.message.tool_calls:
for tool_call in choice.message.tool_calls:
content.append(
{
"type": "function",
"id": tool_call.id,
"function": {
"name": tool_call.function.name,
"arguments": tool_call.function.arguments,
},
}
)
if content:
message = {
"role": role,
"content": content,
}
output.append(message)
# Handle Responses API format
if hasattr(response, "output"):
content = []
role = "assistant"
for item in response.output:
if item.type == "message":
# Extract text content from the content list
role = item.role
if hasattr(item, "content") and isinstance(item.content, list):
for content_item in item.content:
if (
@@ -157,112 +205,110 @@ def format_response_openai(response):
and content_item.type == "output_text"
and hasattr(content_item, "text")
):
output.append(
{
"content": content_item.text,
"role": item.role,
}
)
content.append({"type": "text", "text": content_item.text})
elif hasattr(content_item, "text"):
output.append(
{
"content": content_item.text,
"role": item.role,
}
)
content.append({"type": "text", "text": content_item.text})
elif (
hasattr(content_item, "type")
and content_item.type == "input_image"
and hasattr(content_item, "image_url")
):
output.append(
content.append(
{
"content": {
"type": "image",
"image": content_item.image_url,
},
"role": item.role,
"type": "image",
"image": content_item.image_url,
}
)
else:
output.append(
{
"content": item.content,
"role": item.role,
}
)
elif hasattr(item, "content"):
content.append({"type": "text", "text": str(item.content)})
elif hasattr(item, "type") and item.type == "function_call":
content.append(
{
"type": "function",
"id": getattr(item, "call_id", getattr(item, "id", "")),
"function": {
"name": item.name,
"arguments": getattr(item, "arguments", {}),
},
}
)
if content:
message = {
"role": role,
"content": content,
}
output.append(message)
return output
def format_response_gemini(response):
output = []
if hasattr(response, "candidates") and response.candidates:
for candidate in response.candidates:
if hasattr(candidate, "content") and candidate.content:
content_text = ""
content = []
if hasattr(candidate.content, "parts") and candidate.content.parts:
for part in candidate.content.parts:
if hasattr(part, "text") and part.text:
content_text += part.text
if content_text:
output.append(
{
"role": "assistant",
"content": content_text,
}
)
content.append({"type": "text", "text": part.text})
elif hasattr(part, "function_call") and part.function_call:
function_call = part.function_call
content.append(
{
"type": "function",
"function": {
"name": function_call.name,
"arguments": function_call.args,
},
}
)
if content:
message = {
"role": "assistant",
"content": content,
}
output.append(message)
elif hasattr(candidate, "text") and candidate.text:
output.append(
{
"role": "assistant",
"content": candidate.text,
"content": [{"type": "text", "text": candidate.text}],
}
)
elif hasattr(response, "text") and response.text:
output.append(
{
"role": "assistant",
"content": response.text,
"content": [{"type": "text", "text": response.text}],
}
)
return output
def format_tool_calls(response, provider: str):
def extract_available_tool_calls(provider: str, kwargs: Dict[str, Any]):
if provider == "anthropic":
if hasattr(response, "content") and response.content:
tool_calls = []
if "tools" in kwargs:
return kwargs["tools"]
for content_item in response.content:
if hasattr(content_item, "type") and content_item.type == "tool_use":
tool_calls.append(
{
"type": content_item.type,
"id": content_item.id,
"name": content_item.name,
"input": content_item.input,
}
)
return None
elif provider == "gemini":
if "config" in kwargs and hasattr(kwargs["config"], "tools"):
return kwargs["config"].tools
return tool_calls if tool_calls else None
return None
elif provider == "openai":
# Handle both Chat Completions and Responses API
if hasattr(response, "choices") and response.choices:
# Check for tool_calls in message (Chat Completions format)
if (
hasattr(response.choices[0], "message")
and hasattr(response.choices[0].message, "tool_calls")
and response.choices[0].message.tool_calls
):
return response.choices[0].message.tool_calls
if "tools" in kwargs:
return kwargs["tools"]
# Check for tool_calls directly in response (Responses API format)
if (
hasattr(response.choices[0], "tool_calls")
and response.choices[0].tool_calls
):
return response.choices[0].tool_calls
return None
return None
def merge_system_prompt(kwargs: Dict[str, Any], provider: str):
@@ -395,12 +441,10 @@ def call_llm_and_track_usage(
**(error_params or {}),
}
tool_calls = format_tool_calls(response, provider)
available_tool_calls = extract_available_tool_calls(provider, kwargs)
if tool_calls:
event_properties["$ai_tools"] = with_privacy_mode(
ph_client, posthog_privacy_mode, tool_calls
)
if available_tool_calls:
event_properties["$ai_tools"] = available_tool_calls
if (
usage.get("cache_read_input_tokens") is not None
@@ -511,11 +555,10 @@ async def call_llm_and_track_usage_async(
**(error_params or {}),
}
tool_calls = format_tool_calls(response, provider)
if tool_calls:
event_properties["$ai_tools"] = with_privacy_mode(
ph_client, posthog_privacy_mode, tool_calls
)
available_tool_calls = extract_available_tool_calls(provider, kwargs)
if available_tool_calls:
event_properties["$ai_tools"] = available_tool_calls
if (
usage.get("cache_read_input_tokens") is not None
+53 -33
View File
@@ -83,6 +83,7 @@ def get_identity_state(passed) -> tuple[str, bool]:
def add_context_tags(properties):
properties = properties or {}
current_context = _get_current_context()
if current_context:
context_tags = current_context.collect_tags()
@@ -395,7 +396,7 @@ class Client(object):
def get_flags_decision(
self,
distinct_id: Optional[ID_TYPES] = None,
groups: Optional[dict] = {},
groups: Optional[dict] = None,
person_properties=None,
group_properties=None,
disable_geoip=None,
@@ -418,6 +419,9 @@ class Client(object):
Category:
Feature Flags
"""
groups = groups or {}
person_properties = person_properties or {}
group_properties = group_properties or {}
if distinct_id is None:
distinct_id = get_context_distinct_id()
@@ -505,6 +509,7 @@ class Client(object):
properties = {**(properties or {}), **system_context()}
properties = add_context_tags(properties)
assert properties is not None # Type hint for mypy
(distinct_id, personless) = get_identity_state(distinct_id)
@@ -520,7 +525,7 @@ class Client(object):
}
if groups:
msg["properties"]["$groups"] = groups
properties["$groups"] = groups
extra_properties: dict[str, Any] = {}
feature_variants: Optional[dict[str, Union[bool, str]]] = {}
@@ -575,7 +580,8 @@ class Client(object):
extra_properties["$active_feature_flags"] = active_feature_flags
if extra_properties:
msg["properties"] = {**extra_properties, **msg["properties"]}
properties = {**extra_properties, **properties}
msg["properties"] = properties
return self._enqueue(msg, disable_geoip)
@@ -1153,11 +1159,15 @@ class Client(object):
feature_flag,
distinct_id,
*,
groups={},
person_properties={},
group_properties={},
groups=None,
person_properties=None,
group_properties=None,
warn_on_unknown_groups=True,
) -> FlagValue:
groups = groups or {}
person_properties = person_properties or {}
group_properties = group_properties or {}
if feature_flag.get("ensure_experience_continuity", False):
raise InconclusiveMatchError("Flag has experience continuity enabled")
@@ -1203,9 +1213,9 @@ class Client(object):
key,
distinct_id,
*,
groups={},
person_properties={},
group_properties={},
groups=None,
person_properties=None,
group_properties=None,
only_evaluate_locally=False,
send_feature_flag_events=True,
disable_geoip=None,
@@ -1256,9 +1266,9 @@ class Client(object):
distinct_id: ID_TYPES,
*,
override_match_value: Optional[FlagValue] = None,
groups: Dict[str, str] = {},
person_properties={},
group_properties={},
groups: Optional[Dict[str, str]] = None,
person_properties=None,
group_properties=None,
only_evaluate_locally=False,
send_feature_flag_events=True,
disable_geoip=None,
@@ -1268,9 +1278,16 @@ class Client(object):
person_properties, group_properties = (
self._add_local_person_and_group_properties(
distinct_id, groups, person_properties, group_properties
distinct_id,
groups or {},
person_properties or {},
group_properties or {},
)
)
# Ensure non-None values for type checking
groups = groups or {}
person_properties = person_properties or {}
group_properties = group_properties or {}
flag_result = None
flag_details = None
@@ -1285,7 +1302,7 @@ class Client(object):
lookup_match_value = override_match_value or flag_value
payload = (
self._compute_payload_locally(key, lookup_match_value)
if lookup_match_value
if lookup_match_value is not None
else None
)
flag_result = FeatureFlagResult.from_value_and_payload(
@@ -1354,9 +1371,9 @@ class Client(object):
key,
distinct_id,
*,
groups={},
person_properties={},
group_properties={},
groups=None,
person_properties=None,
group_properties=None,
only_evaluate_locally=False,
send_feature_flag_events=True,
disable_geoip=None,
@@ -1404,9 +1421,9 @@ class Client(object):
key,
distinct_id,
*,
groups={},
person_properties={},
group_properties={},
groups=None,
person_properties=None,
group_properties=None,
only_evaluate_locally=False,
send_feature_flag_events=True,
disable_geoip=None,
@@ -1492,9 +1509,9 @@ class Client(object):
distinct_id,
*,
match_value: Optional[FlagValue] = None,
groups={},
person_properties={},
group_properties={},
groups=None,
person_properties=None,
group_properties=None,
only_evaluate_locally=False,
send_feature_flag_events=True,
disable_geoip=None,
@@ -1586,7 +1603,7 @@ class Client(object):
f"$feature/{key}": response,
}
if payload:
if payload is not None:
# if payload is not a string, json serialize it to a string
properties["$feature_flag_payload"] = payload
@@ -1662,9 +1679,9 @@ class Client(object):
self,
distinct_id,
*,
groups={},
person_properties={},
group_properties={},
groups=None,
person_properties=None,
group_properties=None,
only_evaluate_locally=False,
disable_geoip=None,
) -> Optional[dict[str, Union[bool, str]]]:
@@ -1702,9 +1719,9 @@ class Client(object):
self,
distinct_id,
*,
groups={},
person_properties={},
group_properties={},
groups=None,
person_properties=None,
group_properties=None,
only_evaluate_locally=False,
disable_geoip=None,
) -> FlagsAndPayloads:
@@ -1765,10 +1782,13 @@ class Client(object):
distinct_id: ID_TYPES,
*,
groups: Dict[str, Union[str, int]],
person_properties={},
group_properties={},
person_properties=None,
group_properties=None,
warn_on_unknown_groups=False,
) -> tuple[FlagsAndPayloads, bool]:
person_properties = person_properties or {}
group_properties = group_properties or {}
if self.feature_flags is None and self.personal_api_key:
self.load_feature_flags()
@@ -1790,7 +1810,7 @@ class Client(object):
matched_payload = self._compute_payload_locally(
flag["key"], flags[flag["key"]]
)
if matched_payload:
if matched_payload is not None:
payloads[flag["key"]] = matched_payload
except InconclusiveMatchError:
# No need to log this, since it's just telling us to fall back to `/decide`
+270 -29
View File
@@ -89,26 +89,51 @@ def mock_anthropic_response_with_cached_tokens():
@pytest.fixture
def mock_anthropic_response_with_tool_use():
def mock_anthropic_response_with_tool_calls():
return Message(
id="msg_123",
id="msg_456",
type="message",
role="assistant",
content=[
{"type": "text", "text": "I'll help you with that."},
{"type": "text", "text": "I'll help you check the weather."},
{"type": "text", "text": " Let me look that up."},
{
"type": "tool_use",
"id": "tool_1",
"id": "toolu_abc123",
"name": "get_weather",
"input": {"location": "New York"},
"input": {"location": "San Francisco"},
},
],
model="claude-3-opus-20240229",
model="claude-3-5-sonnet-20241022",
usage=Usage(
input_tokens=20,
output_tokens=10,
input_tokens=25,
output_tokens=15,
),
stop_reason="end_turn",
stop_reason="tool_use",
stop_sequence=None,
)
@pytest.fixture
def mock_anthropic_response_tool_calls_only():
return Message(
id="msg_789",
type="message",
role="assistant",
content=[
{
"type": "tool_use",
"id": "toolu_def456",
"name": "get_weather",
"input": {"location": "New York", "unit": "fahrenheit"},
}
],
model="claude-3-5-sonnet-20241022",
usage=Usage(
input_tokens=30,
output_tokens=12,
),
stop_reason="tool_use",
stop_sequence=None,
)
@@ -137,7 +162,10 @@ def test_basic_completion(mock_client, mock_anthropic_response):
assert props["$ai_model"] == "claude-3-opus-20240229"
assert props["$ai_input"] == [{"role": "user", "content": "Hello"}]
assert props["$ai_output_choices"] == [
{"role": "assistant", "content": "Test response"}
{
"role": "assistant",
"content": [{"type": "text", "text": "Test response"}],
}
]
assert props["$ai_input_tokens"] == 20
assert props["$ai_output_tokens"] == 10
@@ -311,7 +339,9 @@ def test_basic_integration(mock_client):
{"role": "user", "content": "Foo"},
]
assert props["$ai_output_choices"][0]["role"] == "assistant"
assert props["$ai_output_choices"][0]["content"] == "Bar"
assert props["$ai_output_choices"][0]["content"] == [
{"type": "text", "text": "Bar"}
]
assert props["$ai_input_tokens"] == 18
assert props["$ai_output_tokens"] == 1
assert props["$ai_http_status"] == 200
@@ -450,7 +480,10 @@ def test_cached_tokens(mock_client, mock_anthropic_response_with_cached_tokens):
assert props["$ai_model"] == "claude-3-opus-20240229"
assert props["$ai_input"] == [{"role": "user", "content": "Hello"}]
assert props["$ai_output_choices"] == [
{"role": "assistant", "content": "Test response"}
{
"role": "assistant",
"content": [{"type": "text", "text": "Test response"}],
}
]
assert props["$ai_input_tokens"] == 20
assert props["$ai_output_tokens"] == 10
@@ -461,20 +494,41 @@ def test_cached_tokens(mock_client, mock_anthropic_response_with_cached_tokens):
assert isinstance(props["$ai_latency"], float)
def test_tool_use_response(mock_client, mock_anthropic_response_with_tool_use):
def test_tool_definition(mock_client, mock_anthropic_response):
with patch(
"anthropic.resources.Messages.create",
return_value=mock_anthropic_response_with_tool_use,
return_value=mock_anthropic_response,
):
client = Anthropic(api_key="test-key", posthog_client=mock_client)
tools = [
{
"name": "get_weather",
"description": "Get the current weather for a specific location",
"input_schema": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city or location name to get weather for",
}
},
"required": ["location"],
},
}
]
response = client.messages.create(
model="claude-3-opus-20240229",
messages=[{"role": "user", "content": "What's the weather like?"}],
model="claude-3-5-sonnet-20241022",
max_tokens=200,
temperature=0.7,
tools=tools,
messages=[{"role": "user", "content": "hey"}],
posthog_distinct_id="test-id",
posthog_properties={"foo": "bar"},
)
assert response == mock_anthropic_response_with_tool_use
assert response == mock_anthropic_response
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
@@ -483,25 +537,212 @@ def test_tool_use_response(mock_client, mock_anthropic_response_with_tool_use):
assert call_args["distinct_id"] == "test-id"
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "anthropic"
assert props["$ai_model"] == "claude-3-opus-20240229"
assert props["$ai_input"] == [
{"role": "user", "content": "What's the weather like?"}
]
# Should only include text content, not tool_use content
assert props["$ai_model"] == "claude-3-5-sonnet-20241022"
assert props["$ai_input"] == [{"role": "user", "content": "hey"}]
assert props["$ai_output_choices"] == [
{"role": "assistant", "content": "I'll help you with that."}
{
"role": "assistant",
"content": [{"type": "text", "text": "Test response"}],
}
]
assert props["$ai_input_tokens"] == 20
assert props["$ai_output_tokens"] == 10
assert props["$ai_http_status"] == 200
assert props["foo"] == "bar"
assert isinstance(props["$ai_latency"], float)
# Verify that tools are captured separately
assert props["$ai_tools"] == [
# Verify that tools are captured in the $ai_tools property
assert props["$ai_tools"] == tools
def test_tool_calls_in_output_choices(
mock_client, mock_anthropic_response_with_tool_calls
):
with patch(
"anthropic.resources.Messages.create",
return_value=mock_anthropic_response_with_tool_calls,
):
client = Anthropic(api_key="test-key", posthog_client=mock_client)
response = client.messages.create(
model="claude-3-5-sonnet-20241022",
max_tokens=200,
messages=[
{"role": "user", "content": "What's the weather in San Francisco?"}
],
tools=[
{
"name": "get_weather",
"description": "Get weather",
"input_schema": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
}
],
posthog_distinct_id="test-id",
)
assert response == mock_anthropic_response_with_tool_calls
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
props = call_args["properties"]
assert call_args["distinct_id"] == "test-id"
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "anthropic"
assert props["$ai_model"] == "claude-3-5-sonnet-20241022"
assert props["$ai_output_choices"] == [
{
"type": "tool_use",
"id": "tool_1",
"name": "get_weather",
"input": {"location": "New York"},
"role": "assistant",
"content": [
{"type": "text", "text": "I'll help you check the weather."},
{"type": "text", "text": " Let me look that up."},
{
"type": "function",
"id": "toolu_abc123",
"function": {
"name": "get_weather",
"arguments": {"location": "San Francisco"},
},
},
],
}
]
# Check token usage
assert props["$ai_input_tokens"] == 25
assert props["$ai_output_tokens"] == 15
assert props["$ai_http_status"] == 200
def test_tool_calls_only_no_content(
mock_client, mock_anthropic_response_tool_calls_only
):
with patch(
"anthropic.resources.Messages.create",
return_value=mock_anthropic_response_tool_calls_only,
):
client = Anthropic(api_key="test-key", posthog_client=mock_client)
response = client.messages.create(
model="claude-3-5-sonnet-20241022",
max_tokens=200,
messages=[{"role": "user", "content": "Get weather for New York"}],
tools=[
{
"name": "get_weather",
"description": "Get weather",
"input_schema": {
"type": "object",
"properties": {
"location": {"type": "string"},
"unit": {"type": "string"},
},
"required": ["location"],
},
}
],
posthog_distinct_id="test-id",
)
assert response == mock_anthropic_response_tool_calls_only
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
props = call_args["properties"]
assert call_args["distinct_id"] == "test-id"
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "anthropic"
assert props["$ai_model"] == "claude-3-5-sonnet-20241022"
assert props["$ai_output_choices"] == [
{
"role": "assistant",
"content": [
{
"type": "function",
"id": "toolu_def456",
"function": {
"name": "get_weather",
"arguments": {"location": "New York", "unit": "fahrenheit"},
},
}
],
}
]
# Check token usage
assert props["$ai_input_tokens"] == 30
assert props["$ai_output_tokens"] == 12
assert props["$ai_http_status"] == 200
def test_async_tool_calls_in_output_choices(
mock_client, mock_anthropic_response_with_tool_calls
):
import asyncio
async def mock_async_create(**kwargs):
return mock_anthropic_response_with_tool_calls
with patch(
"anthropic.resources.AsyncMessages.create",
side_effect=mock_async_create,
):
async_client = AsyncAnthropic(api_key="test-key", posthog_client=mock_client)
async def run_test():
return await async_client.messages.create(
model="claude-3-5-sonnet-20241022",
max_tokens=200,
messages=[
{"role": "user", "content": "What's the weather in San Francisco?"}
],
tools=[
{
"name": "get_weather",
"description": "Get weather",
"input_schema": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
}
],
posthog_distinct_id="test-id",
)
response = asyncio.run(run_test())
assert response == mock_anthropic_response_with_tool_calls
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
props = call_args["properties"]
assert call_args["distinct_id"] == "test-id"
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "anthropic"
assert props["$ai_model"] == "claude-3-5-sonnet-20241022"
assert props["$ai_output_choices"] == [
{
"role": "assistant",
"content": [
{"type": "text", "text": "I'll help you check the weather."},
{"type": "text", "text": " Let me look that up."},
{
"type": "function",
"id": "toolu_abc123",
"function": {
"name": "get_weather",
"arguments": {"location": "San Francisco"},
},
},
],
}
]
# Check token usage
assert props["$ai_input_tokens"] == 25
assert props["$ai_output_tokens"] == 15
assert props["$ai_http_status"] == 200
+311
View File
@@ -56,6 +56,87 @@ def mock_google_genai_client():
yield mock_client_instance
@pytest.fixture
def mock_gemini_response_with_function_calls():
mock_response = MagicMock()
# Mock usage metadata
mock_usage = MagicMock()
mock_usage.prompt_token_count = 25
mock_usage.candidates_token_count = 15
mock_response.usage_metadata = mock_usage
# Mock function call
mock_function_call = MagicMock()
mock_function_call.name = "get_current_weather"
mock_function_call.args = {"location": "San Francisco"}
# Mock text part 1
mock_text_part1 = MagicMock()
mock_text_part1.text = "I'll check the weather for you."
# Make hasattr(part, "text") return True
type(mock_text_part1).text = mock_text_part1.text
# Mock text part 2
mock_text_part2 = MagicMock()
mock_text_part2.text = " Let me look that up."
type(mock_text_part2).text = mock_text_part2.text
# Mock function call part - need to ensure hasattr() works correctly
mock_function_part = MagicMock()
mock_function_part.function_call = mock_function_call
# Make hasattr(part, "function_call") return True
type(mock_function_part).function_call = mock_function_part.function_call
# Ensure hasattr(part, "text") returns False for the function part
del mock_function_part.text
# Mock content with 2 text parts and 1 function call part
mock_content = MagicMock()
mock_content.parts = [mock_text_part1, mock_text_part2, mock_function_part]
# Mock candidate
mock_candidate = MagicMock()
mock_candidate.content = mock_content
mock_response.candidates = [mock_candidate]
return mock_response
@pytest.fixture
def mock_gemini_response_function_calls_only():
mock_response = MagicMock()
# Mock usage metadata
mock_usage = MagicMock()
mock_usage.prompt_token_count = 30
mock_usage.candidates_token_count = 12
mock_response.usage_metadata = mock_usage
# Mock function call
mock_function_call = MagicMock()
mock_function_call.name = "get_current_weather"
mock_function_call.args = {"location": "New York", "unit": "fahrenheit"}
# Mock function call part (no text part) - need to ensure hasattr() works correctly
mock_function_part = MagicMock()
mock_function_part.function_call = mock_function_call
# Make hasattr(part, "function_call") return True
type(mock_function_part).function_call = mock_function_part.function_call
# Ensure hasattr(part, "text") returns False for the function part
del mock_function_part.text
# Mock content with only function call part
mock_content = MagicMock()
mock_content.parts = [mock_function_part]
# Mock candidate
mock_candidate = MagicMock()
mock_candidate.content = mock_content
mock_response.candidates = [mock_candidate]
return mock_response
def test_new_client_basic_generation(
mock_client, mock_google_genai_client, mock_gemini_response
):
@@ -318,3 +399,233 @@ def test_new_client_override_defaults(
assert props["team"] == "ai" # from defaults
assert props["feature"] == "chat" # from call
assert props["urgent"] is True # from call
def test_vertex_ai_parameters_passed_through(
mock_client, mock_google_genai_client, mock_gemini_response
):
"""Test that Vertex AI parameters are properly passed to genai.Client"""
mock_google_genai_client.models.generate_content.return_value = mock_gemini_response
# Mock credentials object
mock_credentials = MagicMock()
mock_debug_config = MagicMock()
mock_http_options = MagicMock()
# Create client with Vertex AI parameters
Client(
vertexai=True,
credentials=mock_credentials,
project="test-project",
location="us-central1",
debug_config=mock_debug_config,
http_options=mock_http_options,
posthog_client=mock_client,
)
# Verify genai.Client was called with correct parameters
google_genai.Client.assert_called_once_with(
vertexai=True,
credentials=mock_credentials,
project="test-project",
location="us-central1",
debug_config=mock_debug_config,
http_options=mock_http_options,
)
def test_api_key_mode(mock_client, mock_google_genai_client):
"""Test API key authentication mode"""
# Create client with just API key (traditional mode)
Client(
api_key="test-api-key",
posthog_client=mock_client,
)
# Verify genai.Client was called with only api_key
google_genai.Client.assert_called_once_with(api_key="test-api-key")
def test_vertex_ai_mode_with_optional_api_key(
mock_client, mock_google_genai_client, mock_gemini_response
):
"""Test Vertex AI mode with optional API key"""
mock_google_genai_client.models.generate_content.return_value = mock_gemini_response
mock_credentials = MagicMock()
# Create client with Vertex AI + API key
Client(
vertexai=True,
api_key="test-api-key",
credentials=mock_credentials,
project="test-project",
posthog_client=mock_client,
)
# Verify genai.Client was called with both Vertex AI params and API key
google_genai.Client.assert_called_once_with(
vertexai=True,
api_key="test-api-key",
credentials=mock_credentials,
project="test-project",
)
def test_tool_use_response(mock_client, mock_google_genai_client, mock_gemini_response):
"""Test that tools defined in config are captured in $ai_tools property"""
mock_google_genai_client.models.generate_content.return_value = mock_gemini_response
client = Client(api_key="test-key", posthog_client=mock_client)
# Create mock tools configuration
mock_tool = MagicMock()
mock_tool.function_declarations = [
MagicMock(
name="get_current_weather",
description="Gets the current weather for a given location.",
parameters=MagicMock(
type="OBJECT",
properties={
"location": MagicMock(
type="STRING",
description="The city and state, e.g. San Francisco, CA",
)
},
required=["location"],
),
)
]
mock_config = MagicMock()
mock_config.tools = [mock_tool]
response = client.models.generate_content(
model="gemini-2.5-flash",
contents=["hey"],
config=mock_config,
posthog_distinct_id="test-id",
posthog_properties={"foo": "bar"},
)
assert response == mock_gemini_response
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
props = call_args["properties"]
assert call_args["distinct_id"] == "test-id"
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "gemini"
assert props["$ai_model"] == "gemini-2.5-flash"
assert props["$ai_input"] == [{"role": "user", "content": "hey"}]
assert props["$ai_output_choices"] == [
{
"role": "assistant",
"content": [{"type": "text", "text": "Test response from Gemini"}],
}
]
assert props["$ai_input_tokens"] == 20
assert props["$ai_output_tokens"] == 10
assert props["$ai_http_status"] == 200
assert props["foo"] == "bar"
assert isinstance(props["$ai_latency"], float)
# Verify that tools are captured in the $ai_tools property
assert props["$ai_tools"] == [mock_tool]
def test_function_calls_in_output_choices(
mock_client, mock_google_genai_client, mock_gemini_response_with_function_calls
):
"""Test that function calls are properly included in $ai_output_choices"""
mock_google_genai_client.models.generate_content.return_value = (
mock_gemini_response_with_function_calls
)
client = Client(api_key="test-key", posthog_client=mock_client)
response = client.models.generate_content(
model="gemini-2.5-flash",
contents=["What's the weather in San Francisco?"],
posthog_distinct_id="test-id",
)
assert response == mock_gemini_response_with_function_calls
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
props = call_args["properties"]
assert call_args["distinct_id"] == "test-id"
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "gemini"
assert props["$ai_model"] == "gemini-2.5-flash"
assert props["$ai_output_choices"] == [
{
"role": "assistant",
"content": [
{"type": "text", "text": "I'll check the weather for you."},
{"type": "text", "text": " Let me look that up."},
{
"type": "function",
"function": {
"name": "get_current_weather",
"arguments": {"location": "San Francisco"},
},
},
],
}
]
# Check token usage
assert props["$ai_input_tokens"] == 25
assert props["$ai_output_tokens"] == 15
assert props["$ai_http_status"] == 200
def test_function_calls_only_no_content(
mock_client, mock_google_genai_client, mock_gemini_response_function_calls_only
):
"""Test function calls without text content in $ai_output_choices"""
mock_google_genai_client.models.generate_content.return_value = (
mock_gemini_response_function_calls_only
)
client = Client(api_key="test-key", posthog_client=mock_client)
response = client.models.generate_content(
model="gemini-2.5-flash",
contents=["Get weather for New York"],
posthog_distinct_id="test-id",
)
assert response == mock_gemini_response_function_calls_only
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
props = call_args["properties"]
assert call_args["distinct_id"] == "test-id"
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "gemini"
assert props["$ai_model"] == "gemini-2.5-flash"
assert props["$ai_output_choices"] == [
{
"role": "assistant",
"content": [
{
"type": "function",
"function": {
"name": "get_current_weather",
"arguments": {"location": "New York", "unit": "fahrenheit"},
},
}
],
}
]
# Check token usage
assert props["$ai_input_tokens"] == 30
assert props["$ai_output_tokens"] == 12
assert props["$ai_http_status"] == 200
+87 -1
View File
@@ -5,7 +5,7 @@ import os
import time
import uuid
from typing import List, Literal, Optional, TypedDict, Union
from unittest.mock import patch
from unittest.mock import patch, MagicMock
import pytest
@@ -1790,3 +1790,89 @@ def test_convert_message_to_dict_tool_calls():
},
}
]
def test_tool_definition(mock_client):
"""Test that tools defined in invocation parameters are captured in $ai_tools property"""
callbacks = CallbackHandler(mock_client)
run_id = uuid.uuid4()
# Define tools to be passed to the invocation parameters
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather for a specific location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city or location name to get weather for",
}
},
"required": ["location"],
},
},
}
]
with patch("time.time", return_value=1234567890):
callbacks._set_llm_metadata(
{"kwargs": {"openai_api_base": "https://api.openai.com/v1"}},
run_id,
messages=[{"role": "user", "content": "hey"}],
invocation_params={"temperature": 0.7, "tools": tools},
metadata={"ls_model_name": "gpt-4o-mini", "ls_provider": "openai"},
name="test",
)
expected = GenerationMetadata(
model="gpt-4o-mini",
input=[{"role": "user", "content": "hey"}],
start_time=1234567890,
model_params={"temperature": 0.7},
provider="openai",
base_url="https://api.openai.com/v1",
name="test",
tools=tools,
end_time=None,
)
assert callbacks._runs[run_id] == expected
with patch("time.time", return_value=1234567891):
run = callbacks._pop_run_metadata(run_id)
expected.end_time = 1234567891
assert run == expected
assert callbacks._runs == {}
# Now test that the tools are properly captured in the PostHog event
mock_response = MagicMock()
mock_response.generations = [[MagicMock()]]
callbacks._capture_generation(
trace_id=run_id,
run_id=run_id,
run=run,
output=mock_response,
parent_run_id=None,
)
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
props = call_args["properties"]
assert call_args["distinct_id"] == run_id
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "openai"
assert props["$ai_model"] == "gpt-4o-mini"
assert props["$ai_input"] == [{"role": "user", "content": "hey"}]
assert props["$ai_model_parameters"] == {"temperature": 0.7}
assert props["$ai_base_url"] == "https://api.openai.com/v1"
assert props["$ai_span_name"] == "test"
assert props["$ai_span_id"] == run_id
assert props["$ai_trace_id"] == run_id
assert props["$ai_latency"] == 1.0
# Verify that tools are captured in the $ai_tools property
assert props["$ai_tools"] == tools
+335 -34
View File
@@ -1,4 +1,3 @@
import json
import time
from unittest.mock import patch
@@ -26,6 +25,7 @@ try:
ResponseOutputMessage,
ResponseOutputText,
ResponseUsage,
ResponseFunctionToolCall,
ParsedResponse,
)
from openai.types.responses.parsed_response import (
@@ -227,11 +227,27 @@ def mock_openai_response_with_tool_calls():
created=int(time.time()),
choices=[
Choice(
finish_reason="tool_calls",
finish_reason="stop",
index=0,
message=ChatCompletionMessage(
content="I'll check the weather for you.",
role="assistant",
),
),
Choice(
finish_reason="stop",
index=1,
message=ChatCompletionMessage(
content=" Let me look that up.",
role="assistant",
),
),
Choice(
finish_reason="tool_calls",
index=2,
message=ChatCompletionMessage(
content=None,
role="assistant",
tool_calls=[
ChatCompletionMessageToolCall(
id="call_abc123",
@@ -243,7 +259,7 @@ def mock_openai_response_with_tool_calls():
)
],
),
)
),
],
usage=CompletionUsage(
completion_tokens=15,
@@ -253,6 +269,97 @@ def mock_openai_response_with_tool_calls():
)
@pytest.fixture
def mock_openai_response_tool_calls_only():
return ChatCompletion(
id="test",
model="gpt-4",
object="chat.completion",
created=int(time.time()),
choices=[
Choice(
finish_reason="tool_calls",
index=0,
message=ChatCompletionMessage(
content=None,
role="assistant",
tool_calls=[
ChatCompletionMessageToolCall(
id="call_def456",
type="function",
function=Function(
name="get_weather",
arguments='{"location": "New York"}',
),
)
],
),
)
],
usage=CompletionUsage(
completion_tokens=10,
prompt_tokens=25,
total_tokens=35,
),
)
@pytest.fixture
def mock_responses_api_with_tool_calls():
return Response(
id="resp_123",
object="response",
created_at=int(time.time()),
model="gpt-4o-mini",
status="completed",
error=None,
incomplete_details=None,
instructions=None,
max_output_tokens=None,
tools=[],
tool_choice="auto",
parallel_tool_calls=True,
output=[
ResponseOutputMessage(
id="msg_456",
type="message",
role="assistant",
status="completed",
content=[
ResponseOutputText(
type="output_text",
text="I'll help you with the weather.",
annotations=[],
),
ResponseOutputText(
type="output_text",
text=" Let me check that for you.",
annotations=[],
),
],
),
ResponseFunctionToolCall(
id="fc_789",
type="function_call",
name="get_weather",
call_id="call_xyz789",
arguments='{"location": "Chicago"}',
status="completed",
),
],
usage=ResponseUsage(
input_tokens=30,
output_tokens=20,
input_tokens_details={"prompt_tokens": 30, "cached_tokens": 0},
output_tokens_details={"reasoning_tokens": 0},
total_tokens=50,
),
previous_response_id=None,
user=None,
metadata={},
)
def test_basic_completion(mock_client, mock_openai_response):
with patch(
"openai.resources.chat.completions.Completions.create",
@@ -278,7 +385,10 @@ def test_basic_completion(mock_client, mock_openai_response):
assert props["$ai_model"] == "gpt-4"
assert props["$ai_input"] == [{"role": "user", "content": "Hello"}]
assert props["$ai_output_choices"] == [
{"role": "assistant", "content": "Test response"}
{
"role": "assistant",
"content": [{"type": "text", "text": "Test response"}],
}
]
assert props["$ai_input_tokens"] == 20
assert props["$ai_output_tokens"] == 10
@@ -427,7 +537,10 @@ def test_cached_tokens(mock_client, mock_openai_response_with_cached_tokens):
assert props["$ai_model"] == "gpt-4"
assert props["$ai_input"] == [{"role": "user", "content": "Hello"}]
assert props["$ai_output_choices"] == [
{"role": "assistant", "content": "Test response"}
{
"role": "assistant",
"content": [{"type": "text", "text": "Test response"}],
}
]
assert props["$ai_input_tokens"] == 20
assert props["$ai_output_tokens"] == 10
@@ -475,24 +588,34 @@ def test_tool_calls(mock_client, mock_openai_response_with_tool_calls):
{"role": "user", "content": "What's the weather in San Francisco?"}
]
assert props["$ai_output_choices"] == [
{"role": "assistant", "content": "I'll check the weather for you."}
{
"role": "assistant",
"content": [
{"type": "text", "text": "I'll check the weather for you."},
{"type": "text", "text": " Let me look that up."},
{
"type": "function",
"id": "call_abc123",
"function": {
"name": "get_weather",
"arguments": '{"location": "San Francisco", "unit": "celsius"}',
},
},
],
}
]
# Check that tool calls are properly captured
# Check that defined tools are properly captured in $ai_tools
assert "$ai_tools" in props
tool_calls = props["$ai_tools"]
assert len(tool_calls) == 1
defined_tools = props["$ai_tools"]
assert len(defined_tools) == 1
# Verify the tool call details
tool_call = tool_calls[0]
assert tool_call.id == "call_abc123"
assert tool_call.type == "function"
assert tool_call.function.name == "get_weather"
# Verify the arguments
arguments = tool_call.function.arguments
parsed_args = json.loads(arguments)
assert parsed_args == {"location": "San Francisco", "unit": "celsius"}
# Verify the defined tool details
defined_tool = defined_tools[0]
assert defined_tool["type"] == "function"
assert defined_tool["function"]["name"] == "get_weather"
assert defined_tool["function"]["description"] == "Get weather"
assert defined_tool["function"]["parameters"] == {}
# Check token usage
assert props["$ai_input_tokens"] == 20
@@ -500,6 +623,117 @@ def test_tool_calls(mock_client, mock_openai_response_with_tool_calls):
assert props["$ai_http_status"] == 200
def test_tool_calls_only_no_content(mock_client, mock_openai_response_tool_calls_only):
with patch(
"openai.resources.chat.completions.Completions.create",
return_value=mock_openai_response_tool_calls_only,
):
client = OpenAI(api_key="test-key", posthog_client=mock_client)
response = client.chat.completions.create(
model="gpt-4",
messages=[{"role": "user", "content": "Get weather for New York"}],
tools=[
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather",
"parameters": {},
},
}
],
posthog_distinct_id="test-id",
)
assert response == mock_openai_response_tool_calls_only
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
props = call_args["properties"]
assert call_args["distinct_id"] == "test-id"
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "openai"
assert props["$ai_model"] == "gpt-4"
assert props["$ai_output_choices"] == [
{
"role": "assistant",
"content": [
{
"type": "function",
"id": "call_def456",
"function": {
"name": "get_weather",
"arguments": '{"location": "New York"}',
},
}
],
}
]
# Check token usage
assert props["$ai_input_tokens"] == 25
assert props["$ai_output_tokens"] == 10
assert props["$ai_http_status"] == 200
def test_responses_api_tool_calls(mock_client, mock_responses_api_with_tool_calls):
with patch(
"openai.resources.responses.Responses.create",
return_value=mock_responses_api_with_tool_calls,
):
client = OpenAI(api_key="test-key", posthog_client=mock_client)
response = client.responses.create(
model="gpt-4o-mini",
input=[{"role": "user", "content": "What's the weather in Chicago?"}],
tools=[
{
"name": "get_weather",
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather",
"parameters": {},
},
}
],
posthog_distinct_id="test-id",
)
assert response == mock_responses_api_with_tool_calls
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
props = call_args["properties"]
assert call_args["distinct_id"] == "test-id"
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "openai"
assert props["$ai_model"] == "gpt-4o-mini"
assert props["$ai_output_choices"] == [
{
"role": "assistant",
"content": [
{"type": "text", "text": "I'll help you with the weather."},
{"type": "text", "text": " Let me check that for you."},
{
"type": "function",
"id": "call_xyz789",
"function": {
"name": "get_weather",
"arguments": '{"location": "Chicago"}',
},
},
],
}
]
# Check token usage
assert props["$ai_input_tokens"] == 30
assert props["$ai_output_tokens"] == 20
assert props["$ai_http_status"] == 200
def test_streaming_with_tool_calls(mock_client):
# Create mock tool call chunks that will be returned in sequence
tool_call_chunks = [
@@ -644,21 +878,17 @@ def test_streaming_with_tool_calls(mock_client):
assert props["$ai_provider"] == "openai"
assert props["$ai_model"] == "gpt-4"
# Check that the tool calls were properly accumulated
# Check that defined tools are properly captured in $ai_tools
assert "$ai_tools" in props
tool_calls = props["$ai_tools"]
assert len(tool_calls) == 1
defined_tools = props["$ai_tools"]
assert len(defined_tools) == 1
# Verify the complete tool call was properly assembled
tool_call = tool_calls[0]
assert tool_call.id == "call_abc123"
assert tool_call.type == "function"
assert tool_call.function.name == "get_weather"
# Verify the arguments were concatenated correctly
arguments = tool_call.function.arguments
parsed_args = json.loads(arguments)
assert parsed_args == {"location": "San Francisco", "unit": "celsius"}
# Verify the defined tool details
defined_tool = defined_tools[0]
assert defined_tool["type"] == "function"
assert defined_tool["function"]["name"] == "get_weather"
assert defined_tool["function"]["description"] == "Get weather"
assert defined_tool["function"]["parameters"] == {}
# Check that the content was also accumulated
assert (
@@ -696,7 +926,10 @@ def test_responses_api(mock_client, mock_openai_response_with_responses_api):
assert props["$ai_model"] == "gpt-4o-mini"
assert props["$ai_input"] == [{"role": "user", "content": "Hello"}]
assert props["$ai_output_choices"] == [
{"role": "assistant", "content": "Test response"}
{
"role": "assistant",
"content": [{"type": "text", "text": "Test response"}],
}
]
assert props["$ai_input_tokens"] == 10
assert props["$ai_output_tokens"] == 10
@@ -765,7 +998,12 @@ def test_responses_parse(mock_client, mock_parsed_response):
assert props["$ai_output_choices"] == [
{
"role": "assistant",
"content": '{"name": "Science Fair", "date": "Friday", "participants": ["Alice", "Bob"]}',
"content": [
{
"type": "text",
"text": '{"name": "Science Fair", "date": "Friday", "participants": ["Alice", "Bob"]}',
}
],
}
]
assert props["$ai_input_tokens"] == 15
@@ -774,3 +1012,66 @@ def test_responses_parse(mock_client, mock_parsed_response):
assert props["$ai_http_status"] == 200
assert props["foo"] == "bar"
assert isinstance(props["$ai_latency"], float)
def test_tool_definition(mock_client, mock_openai_response):
"""Test that tools defined in the create function are captured in $ai_tools property"""
with patch(
"openai.resources.chat.completions.Completions.create",
return_value=mock_openai_response,
):
client = OpenAI(api_key="test-key", posthog_client=mock_client)
# Define tools to be passed to the create function
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather for a specific location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city or location name to get weather for",
}
},
"required": ["location"],
},
},
}
]
response = client.chat.completions.create(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "hey"}],
tools=tools,
posthog_distinct_id="test-id",
posthog_properties={"foo": "bar"},
)
assert response == mock_openai_response
assert mock_client.capture.call_count == 1
call_args = mock_client.capture.call_args[1]
props = call_args["properties"]
assert call_args["distinct_id"] == "test-id"
assert call_args["event"] == "$ai_generation"
assert props["$ai_provider"] == "openai"
assert props["$ai_model"] == "gpt-4o-mini"
assert props["$ai_input"] == [{"role": "user", "content": "hey"}]
assert props["$ai_output_choices"] == [
{
"role": "assistant",
"content": [{"type": "text", "text": "Test response"}],
}
]
assert props["$ai_input_tokens"] == 20
assert props["$ai_output_tokens"] == 10
assert props["$ai_http_status"] == 200
assert props["foo"] == "bar"
assert isinstance(props["$ai_latency"], float)
# Verify that tools are captured in the $ai_tools property
assert props["$ai_tools"] == tools
+113 -4
View File
@@ -647,8 +647,8 @@ class TestClient(unittest.TestCase):
timeout=3,
distinct_id="distinct_id",
groups={},
person_properties=None,
group_properties=None,
person_properties={},
group_properties={},
geoip_disable=True,
)
@@ -711,8 +711,8 @@ class TestClient(unittest.TestCase):
timeout=12,
distinct_id="distinct_id",
groups={},
person_properties=None,
group_properties=None,
person_properties={},
group_properties={},
geoip_disable=False,
)
@@ -2246,3 +2246,112 @@ class TestClient(unittest.TestCase):
with self.assertRaises(TypeError) as cm:
client._parse_send_feature_flags(None)
self.assertIn("Invalid type for send_feature_flags", str(cm.exception))
@mock.patch("posthog.client.batch_post")
def test_get_feature_flag_result_with_empty_string_payload(self, patch_batch_post):
"""Test that get_feature_flag_result returns a FeatureFlagResult when payload is empty string"""
client = Client(
FAKE_TEST_API_KEY,
personal_api_key="test_personal_api_key",
sync_mode=True,
)
# Set up local evaluation with a flag that has empty string payload
client.feature_flags = [
{
"id": 1,
"name": "Test flag",
"key": "test-flag",
"is_simple_flag": False,
"active": True,
"rollout_percentage": None,
"filters": {
"groups": [
{
"properties": [],
"rollout_percentage": None,
"variant": "empty-variant",
}
],
"multivariate": {
"variants": [
{
"key": "empty-variant",
"name": "Empty Variant",
"rollout_percentage": 100,
}
]
},
"payloads": {
"empty-variant": "" # Empty string payload
},
},
}
]
# Test get_feature_flag_result
result = client.get_feature_flag_result(
"test-flag", "test-user", only_evaluate_locally=True
)
# Should return a FeatureFlagResult, not None
self.assertIsNotNone(result)
self.assertEqual(result.key, "test-flag")
self.assertEqual(result.get_value(), "empty-variant")
self.assertEqual(result.payload, "") # Should be empty string, not None
@mock.patch("posthog.client.batch_post")
def test_get_all_flags_and_payloads_with_empty_string(self, patch_batch_post):
"""Test that get_all_flags_and_payloads includes flags with empty string payloads"""
client = Client(
FAKE_TEST_API_KEY,
personal_api_key="test_personal_api_key",
sync_mode=True,
)
# Set up multiple flags with different payload types
client.feature_flags = [
{
"id": 1,
"name": "Flag with empty payload",
"key": "empty-payload-flag",
"is_simple_flag": False,
"active": True,
"filters": {
"groups": [{"properties": [], "variant": "variant1"}],
"multivariate": {
"variants": [{"key": "variant1", "rollout_percentage": 100}]
},
"payloads": {"variant1": ""}, # Empty string
},
},
{
"id": 2,
"name": "Flag with normal payload",
"key": "normal-payload-flag",
"is_simple_flag": False,
"active": True,
"filters": {
"groups": [{"properties": [], "variant": "variant2"}],
"multivariate": {
"variants": [{"key": "variant2", "rollout_percentage": 100}]
},
"payloads": {"variant2": "normal payload"},
},
},
]
result = client.get_all_flags_and_payloads(
"test-user", only_evaluate_locally=True
)
# Check that both flags are included
self.assertEqual(result["featureFlags"]["empty-payload-flag"], "variant1")
self.assertEqual(result["featureFlags"]["normal-payload-flag"], "variant2")
# Check that empty string payload is included (not filtered out)
self.assertIn("empty-payload-flag", result["featureFlagPayloads"])
self.assertEqual(result["featureFlagPayloads"]["empty-payload-flag"], "")
self.assertEqual(
result["featureFlagPayloads"]["normal-payload-flag"], "normal payload"
)
+8 -3
View File
@@ -110,7 +110,7 @@ class FeatureFlag:
variant=variant,
reason=None,
metadata=LegacyFlagMetadata(
payload=payload if payload else None,
payload=payload,
),
)
@@ -178,7 +178,9 @@ class FeatureFlagResult:
key=key,
enabled=enabled,
variant=variant,
payload=json.loads(payload) if isinstance(payload, str) else payload,
payload=json.loads(payload)
if isinstance(payload, str) and payload
else payload,
reason=None,
)
@@ -219,6 +221,7 @@ class FeatureFlagResult:
payload=(
json.loads(details.metadata.payload)
if isinstance(details.metadata.payload, str)
and details.metadata.payload
else details.metadata.payload
),
reason=details.reason.description if details.reason else None,
@@ -296,5 +299,7 @@ def to_payloads(response: FlagsResponse) -> Optional[dict[str, str]]:
return {
key: value.metadata.payload
for key, value in response.get("flags", {}).items()
if isinstance(value, FeatureFlag) and value.enabled and value.metadata.payload
if isinstance(value, FeatureFlag)
and value.enabled
and value.metadata.payload is not None
}
+1 -1
View File
@@ -1,4 +1,4 @@
VERSION = "6.3.2"
VERSION = "6.4.0"
if __name__ == "__main__":
print(VERSION, end="") # noqa: T201