Compare commits
23
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2f1ac45f08 | ||
|
|
50b0c7170a | ||
|
|
f719c3dadf | ||
|
|
105090a6ba | ||
|
|
edfadcc6a8 | ||
|
|
13184e2e16 | ||
|
|
1b8642331f | ||
|
|
6af129f414 | ||
|
|
02e82a6050 | ||
|
|
9a05db8b20 | ||
|
|
e06830e068 | ||
|
|
465baea6f8 | ||
|
|
2bd6e9eaf1 | ||
|
|
e6fe39a0dd | ||
|
|
67f68c00fe | ||
|
|
6156e51f8f | ||
|
|
461c45772a | ||
|
|
a221bffb52 | ||
|
|
26cfd818af | ||
|
|
e868e23dcb | ||
|
|
0bb6342472 | ||
|
|
d76bfe6e5b | ||
|
|
b3e21c1c0e |
@@ -0,0 +1,36 @@
|
||||
version: 2
|
||||
updates:
|
||||
- package-ecosystem: "pip"
|
||||
directory: "/"
|
||||
schedule:
|
||||
interval: "daily"
|
||||
time: "10:00"
|
||||
timezone: "UTC"
|
||||
groups:
|
||||
ai-providers:
|
||||
patterns:
|
||||
- "openai"
|
||||
- "anthropic"
|
||||
- "google-genai"
|
||||
- "langchain-core"
|
||||
- "langchain-community"
|
||||
- "langchain-openai"
|
||||
- "langchain-anthropic"
|
||||
- "langgraph"
|
||||
allow:
|
||||
- dependency-name: "openai"
|
||||
- dependency-name: "anthropic"
|
||||
- dependency-name: "google-genai"
|
||||
- dependency-name: "langchain-core"
|
||||
- dependency-name: "langchain-community"
|
||||
- dependency-name: "langchain-openai"
|
||||
- dependency-name: "langchain-anthropic"
|
||||
- dependency-name: "langgraph"
|
||||
open-pull-requests-limit: 1
|
||||
reviewers:
|
||||
- "PostHog/team-llm-analytics"
|
||||
# Uncomment below to enable auto-merge for minor updates when CI passes
|
||||
# pull-request-branch-name:
|
||||
# separator: "/"
|
||||
# assignees:
|
||||
# - "PostHog/ai-team"
|
||||
@@ -0,0 +1,48 @@
|
||||
name: "Generate References"
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
docs-generation:
|
||||
name: Generate references
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout the repository
|
||||
uses: actions/checkout@85e6279cec87321a52edac9c87bce653a07cf6c2
|
||||
with:
|
||||
fetch-depth: 0
|
||||
token: ${{ secrets.POSTHOG_BOT_PAT }}
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@8d9ed9ac5c53483de85588cdf95a591a75ab9f55
|
||||
with:
|
||||
python-version: 3.11.11
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@0c5e2b8115b80b4c7c5ddf6ffdd634974642d182 # v5.4.1
|
||||
with:
|
||||
enable-cache: true
|
||||
pyproject-file: 'pyproject.toml'
|
||||
|
||||
- name: Generate references
|
||||
run: |
|
||||
uv run bin/docs generate-references
|
||||
|
||||
- name: Check for changes in references
|
||||
id: changes
|
||||
run: |
|
||||
if [ -n "$(git status --porcelain references/)" ]; then
|
||||
echo "changed=true" >> $GITHUB_OUTPUT
|
||||
echo "New references generated in references directory:"
|
||||
git status --porcelain references/
|
||||
else
|
||||
echo "changed=false" >> $GITHUB_OUTPUT
|
||||
echo "No new references generated in references directory"
|
||||
fi
|
||||
|
||||
- uses: stefanzweifel/git-auto-commit-action@778341af668090896ca464160c2def5d1d1a3eb0
|
||||
if: steps.changes.outputs.changed == 'true'
|
||||
with:
|
||||
commit_message: "Update generated references"
|
||||
file_pattern: references/
|
||||
@@ -20,7 +20,7 @@ jobs:
|
||||
uses: actions/checkout@85e6279cec87321a52edac9c87bce653a07cf6c2
|
||||
with:
|
||||
fetch-depth: 0
|
||||
token: ${{ secrets.POSTHOG_BOT_GITHUB_TOKEN }}
|
||||
token: ${{ secrets.POSTHOG_BOT_PAT }}
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@8d9ed9ac5c53483de85588cdf95a591a75ab9f55
|
||||
@@ -45,7 +45,13 @@ jobs:
|
||||
- name: Create GitHub release
|
||||
uses: actions/create-release@0cb9c9b65d5d1901c1f53e5e66eaf4afd303e70e # v1
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.POSTHOG_BOT_GITHUB_TOKEN }}
|
||||
GITHUB_TOKEN: ${{ secrets.POSTHOG_BOT_PAT }}
|
||||
with:
|
||||
tag_name: v${{ env.REPO_VERSION }}
|
||||
release_name: ${{ env.REPO_VERSION }}
|
||||
|
||||
- name: Dispatch generate-references for posthog-python
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
gh workflow run generate-references.yml --ref master
|
||||
@@ -1,3 +1,44 @@
|
||||
# Unreleased
|
||||
|
||||
- fix(django): Restore process_exception method to capture view and downstream middleware exceptions (fixes #329)
|
||||
|
||||
# 6.7.11 - 2025-10-28
|
||||
|
||||
- feat(ai): Add `$ai_framework` property for framework integrations (e.g. LangChain)
|
||||
|
||||
# 6.7.10 - 2025-10-24
|
||||
|
||||
- fix(django): Make middleware truly hybrid - compatible with both sync (WSGI) and async (ASGI) Django stacks without breaking sync-only deployments
|
||||
|
||||
# 6.7.9 - 2025-10-22
|
||||
|
||||
- fix(flags): multi-condition flags with static cohorts returning wrong variants
|
||||
|
||||
# 6.7.8 - 2025-10-16
|
||||
|
||||
- fix(llma): missing async for OpenAI's streaming implementation
|
||||
|
||||
# 6.7.7 - 2025-10-14
|
||||
|
||||
- fix: remove deprecated attribute $exception_personURL from exception events
|
||||
|
||||
# 6.7.6 - 2025-09-16
|
||||
|
||||
- fix: don't sort condition sets with variant overrides to the top
|
||||
- fix: Prevent core Client methods from raising exceptions
|
||||
|
||||
# 6.7.5 - 2025-09-16
|
||||
|
||||
- feat: Django middleware now supports async request handling.
|
||||
|
||||
# 6.7.4 - 2025-09-05
|
||||
|
||||
- fix: Missing system prompts for some providers
|
||||
|
||||
# 6.7.3 - 2025-09-04
|
||||
|
||||
- fix: missing usage tokens in Gemini
|
||||
|
||||
# 6.7.2 - 2025-09-03
|
||||
|
||||
- fix: tool call results in streaming providers
|
||||
|
||||
@@ -3,6 +3,7 @@ Constants for PostHog Python SDK documentation generation.
|
||||
"""
|
||||
|
||||
from typing import Dict, Union
|
||||
from posthog.version import VERSION
|
||||
|
||||
# Documentation generation metadata
|
||||
DOCUMENTATION_METADATA = {
|
||||
@@ -27,8 +28,9 @@ DOCSTRING_PATTERNS = {
|
||||
|
||||
# Output file configuration
|
||||
OUTPUT_CONFIG: Dict[str, Union[str, int]] = {
|
||||
"output_dir": ".",
|
||||
"filename": "posthog-python-references.json",
|
||||
"output_dir": "./references",
|
||||
"filename": f"posthog-python-references-{VERSION}.json",
|
||||
"filename_latest": "posthog-python-references-latest.json",
|
||||
"indent": 2,
|
||||
}
|
||||
|
||||
|
||||
@@ -460,12 +460,23 @@ if __name__ == "__main__":
|
||||
try:
|
||||
documentation = generate_sdk_documentation()
|
||||
|
||||
# Write to file
|
||||
# Ensure output directory exists
|
||||
output_dir = str(OUTPUT_CONFIG["output_dir"])
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
output_file = os.path.join(
|
||||
str(OUTPUT_CONFIG["output_dir"]), str(OUTPUT_CONFIG["filename"])
|
||||
)
|
||||
output_file_latest = os.path.join(
|
||||
str(OUTPUT_CONFIG["output_dir"]), str(OUTPUT_CONFIG["filename_latest"])
|
||||
)
|
||||
|
||||
# Write to current version
|
||||
with open(output_file, "w") as f:
|
||||
json.dump(documentation, f, indent=int(OUTPUT_CONFIG["indent"]))
|
||||
# Write to latest
|
||||
with open(output_file_latest, "w") as f:
|
||||
json.dump(documentation, f, indent=int(OUTPUT_CONFIG["indent"]))
|
||||
|
||||
print(f"✓ Generated {output_file}")
|
||||
|
||||
|
||||
+91
-2
@@ -2,7 +2,13 @@ import datetime # noqa: F401
|
||||
from typing import Callable, Dict, Optional, Any # noqa: F401
|
||||
from typing_extensions import Unpack
|
||||
|
||||
from posthog.args import OptionalCaptureArgs, OptionalSetArgs, ExceptionArg
|
||||
from posthog.args import (
|
||||
OptionalCaptureArgs,
|
||||
OptionalSetArgs,
|
||||
OptionalCaptureAIArgs,
|
||||
AI_EVENT_TYPE,
|
||||
ExceptionArg,
|
||||
)
|
||||
from posthog.client import Client
|
||||
from posthog.contexts import (
|
||||
new_context as inner_new_context,
|
||||
@@ -11,7 +17,15 @@ from posthog.contexts import (
|
||||
set_context_session as inner_set_context_session,
|
||||
identify_context as inner_identify_context,
|
||||
)
|
||||
from posthog.types import FeatureFlag, FlagsAndPayloads, FeatureFlagResult
|
||||
from posthog.feature_flags import (
|
||||
InconclusiveMatchError as InconclusiveMatchError,
|
||||
RequiresServerEvaluation as RequiresServerEvaluation,
|
||||
)
|
||||
from posthog.types import (
|
||||
FeatureFlag,
|
||||
FlagsAndPayloads,
|
||||
FeatureFlagResult as FeatureFlagResult,
|
||||
)
|
||||
from posthog.version import VERSION
|
||||
|
||||
__version__ = VERSION
|
||||
@@ -385,6 +399,81 @@ def capture_exception(
|
||||
return _proxy("capture_exception", exception=exception, **kwargs)
|
||||
|
||||
|
||||
def capture_ai(
|
||||
event: AI_EVENT_TYPE, **kwargs: Unpack[OptionalCaptureAIArgs]
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Capture an AI event to the dedicated AI endpoint with support for large payloads.
|
||||
|
||||
This method sends AI events (like $ai_generation, $ai_trace, etc.) to PostHog's
|
||||
specialized AI endpoint (/i/v0/ai) which supports large payloads through multipart/form-data
|
||||
and blob storage in S3.
|
||||
|
||||
Args:
|
||||
event: The AI event type. Must be one of: "$ai_generation", "$ai_trace", "$ai_span",
|
||||
"$ai_embedding", "$ai_metric", "$ai_feedback"
|
||||
distinct_id: The distinct ID of the user.
|
||||
properties: A dictionary of AI event properties. Must include required properties based on event type:
|
||||
- All events: "$ai_model" (required)
|
||||
- $ai_generation: "$ai_provider", "$ai_trace_id" (required)
|
||||
- $ai_trace: "$ai_trace_id" (required)
|
||||
- $ai_span: "$ai_trace_id", "$ai_span_id" (required)
|
||||
- $ai_embedding: "$ai_provider", "$ai_trace_id" (required)
|
||||
blob_properties: List of property names to send as blobs (large data stored in S3).
|
||||
Common blob properties: "$ai_input", "$ai_output_choices", "$ai_input_state", "$ai_output_state"
|
||||
If not provided, defaults to common blob properties based on event type.
|
||||
timestamp: The timestamp of the event.
|
||||
uuid: A unique identifier for the event.
|
||||
groups: A dictionary of group information.
|
||||
disable_geoip: Whether to disable GeoIP for this event.
|
||||
|
||||
Examples:
|
||||
```python
|
||||
# $ai_generation event with blobs
|
||||
from posthog import capture_ai
|
||||
capture_ai(
|
||||
"$ai_generation",
|
||||
distinct_id="user_123",
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_provider": "openai",
|
||||
"$ai_trace_id": "trace_abc123",
|
||||
"$ai_input": {
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello!"}
|
||||
]
|
||||
},
|
||||
"$ai_output_choices": {
|
||||
"choices": [{"message": {"role": "assistant", "content": "Hi there!"}}]
|
||||
},
|
||||
"$ai_completion_tokens": 150,
|
||||
"$ai_prompt_tokens": 50
|
||||
},
|
||||
blob_properties=["$ai_input", "$ai_output_choices"]
|
||||
)
|
||||
```
|
||||
|
||||
```python
|
||||
# $ai_trace event
|
||||
from posthog import capture_ai
|
||||
capture_ai(
|
||||
"$ai_trace",
|
||||
distinct_id="user_123",
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_trace_id": "trace_abc123"
|
||||
}
|
||||
)
|
||||
```
|
||||
|
||||
Category:
|
||||
AI Events
|
||||
|
||||
Note: This method sends events synchronously to the AI endpoint, bypassing the queue system.
|
||||
"""
|
||||
return _proxy("capture_ai", event, **kwargs)
|
||||
|
||||
|
||||
def feature_enabled(
|
||||
key, # type: str
|
||||
distinct_id, # type: str
|
||||
|
||||
@@ -10,7 +10,7 @@ import time
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from posthog.ai.types import StreamingContentBlock, ToolInProgress
|
||||
from posthog.ai.types import StreamingContentBlock, TokenUsage, ToolInProgress
|
||||
from posthog.ai.utils import (
|
||||
call_llm_and_track_usage,
|
||||
merge_usage_stats,
|
||||
@@ -126,7 +126,7 @@ class WrappedMessages(Messages):
|
||||
**kwargs: Any,
|
||||
):
|
||||
start_time = time.time()
|
||||
usage_stats: Dict[str, int] = {"input_tokens": 0, "output_tokens": 0}
|
||||
usage_stats: TokenUsage = TokenUsage(input_tokens=0, output_tokens=0)
|
||||
accumulated_content = ""
|
||||
content_blocks: List[StreamingContentBlock] = []
|
||||
tools_in_progress: Dict[str, ToolInProgress] = {}
|
||||
@@ -210,14 +210,13 @@ class WrappedMessages(Messages):
|
||||
posthog_privacy_mode: bool,
|
||||
posthog_groups: Optional[Dict[str, Any]],
|
||||
kwargs: Dict[str, Any],
|
||||
usage_stats: Dict[str, int],
|
||||
usage_stats: TokenUsage,
|
||||
latency: float,
|
||||
content_blocks: List[StreamingContentBlock],
|
||||
accumulated_content: str,
|
||||
):
|
||||
from posthog.ai.types import StreamingEventData
|
||||
from posthog.ai.anthropic.anthropic_converter import (
|
||||
standardize_anthropic_usage,
|
||||
format_anthropic_streaming_input,
|
||||
format_anthropic_streaming_output_complete,
|
||||
)
|
||||
@@ -236,7 +235,7 @@ class WrappedMessages(Messages):
|
||||
formatted_output=format_anthropic_streaming_output_complete(
|
||||
content_blocks, accumulated_content
|
||||
),
|
||||
usage_stats=standardize_anthropic_usage(usage_stats),
|
||||
usage_stats=usage_stats,
|
||||
latency=latency,
|
||||
distinct_id=posthog_distinct_id,
|
||||
trace_id=posthog_trace_id,
|
||||
|
||||
@@ -11,7 +11,7 @@ import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from posthog import setup
|
||||
from posthog.ai.types import StreamingContentBlock, ToolInProgress
|
||||
from posthog.ai.types import StreamingContentBlock, TokenUsage, ToolInProgress
|
||||
from posthog.ai.utils import (
|
||||
call_llm_and_track_usage_async,
|
||||
extract_available_tool_calls,
|
||||
@@ -131,7 +131,7 @@ class AsyncWrappedMessages(AsyncMessages):
|
||||
**kwargs: Any,
|
||||
):
|
||||
start_time = time.time()
|
||||
usage_stats: Dict[str, int] = {"input_tokens": 0, "output_tokens": 0}
|
||||
usage_stats: TokenUsage = TokenUsage(input_tokens=0, output_tokens=0)
|
||||
accumulated_content = ""
|
||||
content_blocks: List[StreamingContentBlock] = []
|
||||
tools_in_progress: Dict[str, ToolInProgress] = {}
|
||||
@@ -215,7 +215,7 @@ class AsyncWrappedMessages(AsyncMessages):
|
||||
posthog_privacy_mode: bool,
|
||||
posthog_groups: Optional[Dict[str, Any]],
|
||||
kwargs: Dict[str, Any],
|
||||
usage_stats: Dict[str, int],
|
||||
usage_stats: TokenUsage,
|
||||
latency: float,
|
||||
content_blocks: List[StreamingContentBlock],
|
||||
accumulated_content: str,
|
||||
|
||||
@@ -14,7 +14,6 @@ from posthog.ai.types import (
|
||||
FormattedMessage,
|
||||
FormattedTextContent,
|
||||
StreamingContentBlock,
|
||||
StreamingUsageStats,
|
||||
TokenUsage,
|
||||
ToolInProgress,
|
||||
)
|
||||
@@ -164,7 +163,38 @@ def format_anthropic_streaming_content(
|
||||
return formatted
|
||||
|
||||
|
||||
def extract_anthropic_usage_from_event(event: Any) -> StreamingUsageStats:
|
||||
def extract_anthropic_usage_from_response(response: Any) -> TokenUsage:
|
||||
"""
|
||||
Extract usage from a full Anthropic response (non-streaming).
|
||||
|
||||
Args:
|
||||
response: The complete response from Anthropic API
|
||||
|
||||
Returns:
|
||||
TokenUsage with standardized usage
|
||||
"""
|
||||
if not hasattr(response, "usage"):
|
||||
return TokenUsage(input_tokens=0, output_tokens=0)
|
||||
|
||||
result = TokenUsage(
|
||||
input_tokens=getattr(response.usage, "input_tokens", 0),
|
||||
output_tokens=getattr(response.usage, "output_tokens", 0),
|
||||
)
|
||||
|
||||
if hasattr(response.usage, "cache_read_input_tokens"):
|
||||
cache_read = response.usage.cache_read_input_tokens
|
||||
if cache_read and cache_read > 0:
|
||||
result["cache_read_input_tokens"] = cache_read
|
||||
|
||||
if hasattr(response.usage, "cache_creation_input_tokens"):
|
||||
cache_creation = response.usage.cache_creation_input_tokens
|
||||
if cache_creation and cache_creation > 0:
|
||||
result["cache_creation_input_tokens"] = cache_creation
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def extract_anthropic_usage_from_event(event: Any) -> TokenUsage:
|
||||
"""
|
||||
Extract usage statistics from an Anthropic streaming event.
|
||||
|
||||
@@ -175,7 +205,7 @@ def extract_anthropic_usage_from_event(event: Any) -> StreamingUsageStats:
|
||||
Dictionary of usage statistics
|
||||
"""
|
||||
|
||||
usage: StreamingUsageStats = {}
|
||||
usage: TokenUsage = TokenUsage()
|
||||
|
||||
# Handle usage stats from message_start event
|
||||
if hasattr(event, "type") and event.type == "message_start":
|
||||
@@ -329,26 +359,6 @@ def finalize_anthropic_tool_input(
|
||||
del tools_in_progress[block["id"]]
|
||||
|
||||
|
||||
def standardize_anthropic_usage(usage: Dict[str, Any]) -> TokenUsage:
|
||||
"""
|
||||
Standardize Anthropic usage statistics to common TokenUsage format.
|
||||
|
||||
Anthropic already uses standard field names, so this mainly structures the data.
|
||||
|
||||
Args:
|
||||
usage: Raw usage statistics from Anthropic
|
||||
|
||||
Returns:
|
||||
Standardized TokenUsage dict
|
||||
"""
|
||||
return TokenUsage(
|
||||
input_tokens=usage.get("input_tokens", 0),
|
||||
output_tokens=usage.get("output_tokens", 0),
|
||||
cache_read_input_tokens=usage.get("cache_read_input_tokens"),
|
||||
cache_creation_input_tokens=usage.get("cache_creation_input_tokens"),
|
||||
)
|
||||
|
||||
|
||||
def format_anthropic_streaming_input(kwargs: Dict[str, Any]) -> Any:
|
||||
"""
|
||||
Format Anthropic streaming input using system prompt merging.
|
||||
|
||||
+11
-10
@@ -3,6 +3,9 @@ import time
|
||||
import uuid
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from posthog.ai.types import TokenUsage, StreamingEventData
|
||||
from posthog.ai.utils import merge_system_prompt
|
||||
|
||||
try:
|
||||
from google import genai
|
||||
except ImportError:
|
||||
@@ -17,7 +20,6 @@ from posthog.ai.utils import (
|
||||
merge_usage_stats,
|
||||
)
|
||||
from posthog.ai.gemini.gemini_converter import (
|
||||
format_gemini_input,
|
||||
extract_gemini_usage_from_chunk,
|
||||
extract_gemini_content_from_chunk,
|
||||
format_gemini_streaming_output,
|
||||
@@ -294,7 +296,7 @@ class Models:
|
||||
**kwargs: Any,
|
||||
):
|
||||
start_time = time.time()
|
||||
usage_stats: Dict[str, int] = {"input_tokens": 0, "output_tokens": 0}
|
||||
usage_stats: TokenUsage = TokenUsage(input_tokens=0, output_tokens=0)
|
||||
accumulated_content = []
|
||||
|
||||
kwargs_without_stream = {"model": model, "contents": contents, **kwargs}
|
||||
@@ -350,15 +352,12 @@ class Models:
|
||||
privacy_mode: bool,
|
||||
groups: Optional[Dict[str, Any]],
|
||||
kwargs: Dict[str, Any],
|
||||
usage_stats: Dict[str, int],
|
||||
usage_stats: TokenUsage,
|
||||
latency: float,
|
||||
output: Any,
|
||||
):
|
||||
from posthog.ai.types import StreamingEventData
|
||||
from posthog.ai.gemini.gemini_converter import standardize_gemini_usage
|
||||
|
||||
# Prepare standardized event data
|
||||
formatted_input = self._format_input(contents)
|
||||
formatted_input = self._format_input(contents, **kwargs)
|
||||
sanitized_input = sanitize_gemini(formatted_input)
|
||||
|
||||
event_data = StreamingEventData(
|
||||
@@ -368,7 +367,7 @@ class Models:
|
||||
kwargs=kwargs,
|
||||
formatted_input=sanitized_input,
|
||||
formatted_output=format_gemini_streaming_output(output),
|
||||
usage_stats=standardize_gemini_usage(usage_stats),
|
||||
usage_stats=usage_stats,
|
||||
latency=latency,
|
||||
distinct_id=distinct_id,
|
||||
trace_id=trace_id,
|
||||
@@ -380,10 +379,12 @@ class Models:
|
||||
# Use the common capture function
|
||||
capture_streaming_event(self._ph_client, event_data)
|
||||
|
||||
def _format_input(self, contents):
|
||||
def _format_input(self, contents, **kwargs):
|
||||
"""Format input contents for PostHog tracking"""
|
||||
|
||||
return format_gemini_input(contents)
|
||||
# Create kwargs dict with contents for merge_system_prompt
|
||||
input_kwargs = {"contents": contents, **kwargs}
|
||||
return merge_system_prompt(input_kwargs, "gemini")
|
||||
|
||||
def generate_content_stream(
|
||||
self,
|
||||
|
||||
@@ -10,7 +10,6 @@ from typing import Any, Dict, List, Optional, TypedDict, Union
|
||||
from posthog.ai.types import (
|
||||
FormattedContentItem,
|
||||
FormattedMessage,
|
||||
StreamingUsageStats,
|
||||
TokenUsage,
|
||||
)
|
||||
|
||||
@@ -221,6 +220,30 @@ def format_gemini_response(response: Any) -> List[FormattedMessage]:
|
||||
return output
|
||||
|
||||
|
||||
def extract_gemini_system_instruction(config: Any) -> Optional[str]:
|
||||
"""
|
||||
Extract system instruction from Gemini config parameter.
|
||||
|
||||
Args:
|
||||
config: Config object or dict that may contain system instruction
|
||||
|
||||
Returns:
|
||||
System instruction string if present, None otherwise
|
||||
"""
|
||||
if config is None:
|
||||
return None
|
||||
|
||||
# Handle different config formats
|
||||
if hasattr(config, "system_instruction"):
|
||||
return config.system_instruction
|
||||
elif isinstance(config, dict) and "system_instruction" in config:
|
||||
return config["system_instruction"]
|
||||
elif isinstance(config, dict) and "systemInstruction" in config:
|
||||
return config["systemInstruction"]
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def extract_gemini_tools(kwargs: Dict[str, Any]) -> Optional[Any]:
|
||||
"""
|
||||
Extract tool definitions from Gemini API kwargs.
|
||||
@@ -238,6 +261,38 @@ def extract_gemini_tools(kwargs: Dict[str, Any]) -> Optional[Any]:
|
||||
return None
|
||||
|
||||
|
||||
def format_gemini_input_with_system(
|
||||
contents: Any, config: Any = None
|
||||
) -> List[FormattedMessage]:
|
||||
"""
|
||||
Format Gemini input contents into standardized message format, including system instruction handling.
|
||||
|
||||
Args:
|
||||
contents: Input contents in various possible formats
|
||||
config: Config object or dict that may contain system instruction
|
||||
|
||||
Returns:
|
||||
List of formatted messages with role and content fields, with system message prepended if needed
|
||||
"""
|
||||
formatted_messages = format_gemini_input(contents)
|
||||
|
||||
# Check if system instruction is provided in config parameter
|
||||
system_instruction = extract_gemini_system_instruction(config)
|
||||
|
||||
if system_instruction is not None:
|
||||
has_system = any(msg.get("role") == "system" for msg in formatted_messages)
|
||||
if not has_system:
|
||||
from posthog.ai.types import FormattedMessage
|
||||
|
||||
system_message: FormattedMessage = {
|
||||
"role": "system",
|
||||
"content": system_instruction,
|
||||
}
|
||||
formatted_messages = [system_message] + list(formatted_messages)
|
||||
|
||||
return formatted_messages
|
||||
|
||||
|
||||
def format_gemini_input(contents: Any) -> List[FormattedMessage]:
|
||||
"""
|
||||
Format Gemini input contents into standardized message format for PostHog tracking.
|
||||
@@ -283,7 +338,54 @@ def format_gemini_input(contents: Any) -> List[FormattedMessage]:
|
||||
return [_format_object_message(contents)]
|
||||
|
||||
|
||||
def extract_gemini_usage_from_chunk(chunk: Any) -> StreamingUsageStats:
|
||||
def _extract_usage_from_metadata(metadata: Any) -> TokenUsage:
|
||||
"""
|
||||
Common logic to extract usage from Gemini metadata.
|
||||
Used by both streaming and non-streaming paths.
|
||||
|
||||
Args:
|
||||
metadata: usage_metadata from Gemini response or chunk
|
||||
|
||||
Returns:
|
||||
TokenUsage with standardized usage
|
||||
"""
|
||||
usage = TokenUsage(
|
||||
input_tokens=getattr(metadata, "prompt_token_count", 0),
|
||||
output_tokens=getattr(metadata, "candidates_token_count", 0),
|
||||
)
|
||||
|
||||
# Add cache tokens if present (don't add if 0)
|
||||
if hasattr(metadata, "cached_content_token_count"):
|
||||
cache_tokens = metadata.cached_content_token_count
|
||||
if cache_tokens and cache_tokens > 0:
|
||||
usage["cache_read_input_tokens"] = cache_tokens
|
||||
|
||||
# Add reasoning tokens if present (don't add if 0)
|
||||
if hasattr(metadata, "thoughts_token_count"):
|
||||
reasoning_tokens = metadata.thoughts_token_count
|
||||
if reasoning_tokens and reasoning_tokens > 0:
|
||||
usage["reasoning_tokens"] = reasoning_tokens
|
||||
|
||||
return usage
|
||||
|
||||
|
||||
def extract_gemini_usage_from_response(response: Any) -> TokenUsage:
|
||||
"""
|
||||
Extract usage statistics from a full Gemini response (non-streaming).
|
||||
|
||||
Args:
|
||||
response: The complete response from Gemini API
|
||||
|
||||
Returns:
|
||||
TokenUsage with standardized usage statistics
|
||||
"""
|
||||
if not hasattr(response, "usage_metadata") or not response.usage_metadata:
|
||||
return TokenUsage(input_tokens=0, output_tokens=0)
|
||||
|
||||
return _extract_usage_from_metadata(response.usage_metadata)
|
||||
|
||||
|
||||
def extract_gemini_usage_from_chunk(chunk: Any) -> TokenUsage:
|
||||
"""
|
||||
Extract usage statistics from a Gemini streaming chunk.
|
||||
|
||||
@@ -291,21 +393,16 @@ def extract_gemini_usage_from_chunk(chunk: Any) -> StreamingUsageStats:
|
||||
chunk: Streaming chunk from Gemini API
|
||||
|
||||
Returns:
|
||||
Dictionary of usage statistics
|
||||
TokenUsage with standardized usage statistics
|
||||
"""
|
||||
|
||||
usage: StreamingUsageStats = {}
|
||||
usage: TokenUsage = TokenUsage()
|
||||
|
||||
if not hasattr(chunk, "usage_metadata") or not chunk.usage_metadata:
|
||||
return usage
|
||||
|
||||
# Gemini uses prompt_token_count and candidates_token_count
|
||||
usage["input_tokens"] = getattr(chunk.usage_metadata, "prompt_token_count", 0)
|
||||
usage["output_tokens"] = getattr(chunk.usage_metadata, "candidates_token_count", 0)
|
||||
|
||||
# Calculate total if both values are defined (including 0)
|
||||
if "input_tokens" in usage and "output_tokens" in usage:
|
||||
usage["total_tokens"] = usage["input_tokens"] + usage["output_tokens"]
|
||||
# Use the shared helper to extract usage
|
||||
usage = _extract_usage_from_metadata(chunk.usage_metadata)
|
||||
|
||||
return usage
|
||||
|
||||
@@ -417,22 +514,3 @@ def format_gemini_streaming_output(
|
||||
|
||||
# Fallback for empty or unexpected input
|
||||
return [{"role": "assistant", "content": [{"type": "text", "text": ""}]}]
|
||||
|
||||
|
||||
def standardize_gemini_usage(usage: Dict[str, Any]) -> TokenUsage:
|
||||
"""
|
||||
Standardize Gemini usage statistics to common TokenUsage format.
|
||||
|
||||
Gemini already uses standard field names (input_tokens/output_tokens).
|
||||
|
||||
Args:
|
||||
usage: Raw usage statistics from Gemini
|
||||
|
||||
Returns:
|
||||
Standardized TokenUsage dict
|
||||
"""
|
||||
return TokenUsage(
|
||||
input_tokens=usage.get("input_tokens", 0),
|
||||
output_tokens=usage.get("output_tokens", 0),
|
||||
# Gemini doesn't currently support cache or reasoning tokens
|
||||
)
|
||||
|
||||
@@ -486,6 +486,7 @@ class CallbackHandler(BaseCallbackHandler):
|
||||
"$ai_latency": run.latency,
|
||||
"$ai_span_name": run.name,
|
||||
"$ai_span_id": run_id,
|
||||
"$ai_framework": "langchain",
|
||||
}
|
||||
if parent_run_id is not None:
|
||||
event_properties["$ai_parent_id"] = parent_run_id
|
||||
@@ -556,6 +557,7 @@ class CallbackHandler(BaseCallbackHandler):
|
||||
"$ai_http_status": 200,
|
||||
"$ai_latency": run.latency,
|
||||
"$ai_base_url": run.base_url,
|
||||
"$ai_framework": "langchain",
|
||||
}
|
||||
|
||||
if run.tools:
|
||||
|
||||
@@ -2,6 +2,8 @@ import time
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from posthog.ai.types import TokenUsage
|
||||
|
||||
try:
|
||||
import openai
|
||||
except ImportError:
|
||||
@@ -120,7 +122,7 @@ class WrappedResponses:
|
||||
**kwargs: Any,
|
||||
):
|
||||
start_time = time.time()
|
||||
usage_stats: Dict[str, int] = {}
|
||||
usage_stats: TokenUsage = TokenUsage()
|
||||
final_content = []
|
||||
response = self._original.create(**kwargs)
|
||||
|
||||
@@ -171,14 +173,13 @@ class WrappedResponses:
|
||||
posthog_privacy_mode: bool,
|
||||
posthog_groups: Optional[Dict[str, Any]],
|
||||
kwargs: Dict[str, Any],
|
||||
usage_stats: Dict[str, int],
|
||||
usage_stats: TokenUsage,
|
||||
latency: float,
|
||||
output: Any,
|
||||
available_tool_calls: Optional[List[Dict[str, Any]]] = None,
|
||||
):
|
||||
from posthog.ai.types import StreamingEventData
|
||||
from posthog.ai.openai.openai_converter import (
|
||||
standardize_openai_usage,
|
||||
format_openai_streaming_input,
|
||||
format_openai_streaming_output,
|
||||
)
|
||||
@@ -195,7 +196,7 @@ class WrappedResponses:
|
||||
kwargs=kwargs,
|
||||
formatted_input=sanitized_input,
|
||||
formatted_output=format_openai_streaming_output(output, "responses"),
|
||||
usage_stats=standardize_openai_usage(usage_stats, "responses"),
|
||||
usage_stats=usage_stats,
|
||||
latency=latency,
|
||||
distinct_id=posthog_distinct_id,
|
||||
trace_id=posthog_trace_id,
|
||||
@@ -316,7 +317,7 @@ class WrappedCompletions:
|
||||
**kwargs: Any,
|
||||
):
|
||||
start_time = time.time()
|
||||
usage_stats: Dict[str, int] = {}
|
||||
usage_stats: TokenUsage = TokenUsage()
|
||||
accumulated_content = []
|
||||
accumulated_tool_calls: Dict[int, Dict[str, Any]] = {}
|
||||
if "stream_options" not in kwargs:
|
||||
@@ -387,7 +388,7 @@ class WrappedCompletions:
|
||||
posthog_privacy_mode: bool,
|
||||
posthog_groups: Optional[Dict[str, Any]],
|
||||
kwargs: Dict[str, Any],
|
||||
usage_stats: Dict[str, int],
|
||||
usage_stats: TokenUsage,
|
||||
latency: float,
|
||||
output: Any,
|
||||
tool_calls: Optional[List[Dict[str, Any]]] = None,
|
||||
@@ -395,7 +396,6 @@ class WrappedCompletions:
|
||||
):
|
||||
from posthog.ai.types import StreamingEventData
|
||||
from posthog.ai.openai.openai_converter import (
|
||||
standardize_openai_usage,
|
||||
format_openai_streaming_input,
|
||||
format_openai_streaming_output,
|
||||
)
|
||||
@@ -412,7 +412,7 @@ class WrappedCompletions:
|
||||
kwargs=kwargs,
|
||||
formatted_input=sanitized_input,
|
||||
formatted_output=format_openai_streaming_output(output, "chat", tool_calls),
|
||||
usage_stats=standardize_openai_usage(usage_stats, "chat"),
|
||||
usage_stats=usage_stats,
|
||||
latency=latency,
|
||||
distinct_id=posthog_distinct_id,
|
||||
trace_id=posthog_trace_id,
|
||||
|
||||
@@ -2,6 +2,8 @@ import time
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from posthog.ai.types import TokenUsage
|
||||
|
||||
try:
|
||||
import openai
|
||||
except ImportError:
|
||||
@@ -124,9 +126,9 @@ class WrappedResponses:
|
||||
**kwargs: Any,
|
||||
):
|
||||
start_time = time.time()
|
||||
usage_stats: Dict[str, int] = {}
|
||||
usage_stats: TokenUsage = TokenUsage()
|
||||
final_content = []
|
||||
response = self._original.create(**kwargs)
|
||||
response = await self._original.create(**kwargs)
|
||||
|
||||
async def async_generator():
|
||||
nonlocal usage_stats
|
||||
@@ -176,7 +178,7 @@ class WrappedResponses:
|
||||
posthog_privacy_mode: bool,
|
||||
posthog_groups: Optional[Dict[str, Any]],
|
||||
kwargs: Dict[str, Any],
|
||||
usage_stats: Dict[str, int],
|
||||
usage_stats: TokenUsage,
|
||||
latency: float,
|
||||
output: Any,
|
||||
available_tool_calls: Optional[List[Dict[str, Any]]] = None,
|
||||
@@ -336,14 +338,14 @@ class WrappedCompletions:
|
||||
**kwargs: Any,
|
||||
):
|
||||
start_time = time.time()
|
||||
usage_stats: Dict[str, int] = {}
|
||||
usage_stats: TokenUsage = TokenUsage()
|
||||
accumulated_content = []
|
||||
accumulated_tool_calls: Dict[int, Dict[str, Any]] = {}
|
||||
|
||||
if "stream_options" not in kwargs:
|
||||
kwargs["stream_options"] = {}
|
||||
kwargs["stream_options"]["include_usage"] = True
|
||||
response = self._original.create(**kwargs)
|
||||
response = await self._original.create(**kwargs)
|
||||
|
||||
async def async_generator():
|
||||
nonlocal usage_stats
|
||||
@@ -406,7 +408,7 @@ class WrappedCompletions:
|
||||
posthog_privacy_mode: bool,
|
||||
posthog_groups: Optional[Dict[str, Any]],
|
||||
kwargs: Dict[str, Any],
|
||||
usage_stats: Dict[str, int],
|
||||
usage_stats: TokenUsage,
|
||||
latency: float,
|
||||
output: Any,
|
||||
tool_calls: Optional[List[Dict[str, Any]]] = None,
|
||||
@@ -430,8 +432,8 @@ class WrappedCompletions:
|
||||
format_openai_streaming_output(output, "chat", tool_calls),
|
||||
),
|
||||
"$ai_http_status": 200,
|
||||
"$ai_input_tokens": usage_stats.get("prompt_tokens", 0),
|
||||
"$ai_output_tokens": usage_stats.get("completion_tokens", 0),
|
||||
"$ai_input_tokens": usage_stats.get("input_tokens", 0),
|
||||
"$ai_output_tokens": usage_stats.get("output_tokens", 0),
|
||||
"$ai_cache_read_input_tokens": usage_stats.get(
|
||||
"cache_read_input_tokens", 0
|
||||
),
|
||||
@@ -497,17 +499,17 @@ class WrappedEmbeddings:
|
||||
posthog_trace_id = str(uuid.uuid4())
|
||||
|
||||
start_time = time.time()
|
||||
response = self._original.create(**kwargs)
|
||||
response = await self._original.create(**kwargs)
|
||||
end_time = time.time()
|
||||
|
||||
# Extract usage statistics if available
|
||||
usage_stats = {}
|
||||
usage_stats: TokenUsage = TokenUsage()
|
||||
|
||||
if hasattr(response, "usage") and response.usage:
|
||||
usage_stats = {
|
||||
"prompt_tokens": getattr(response.usage, "prompt_tokens", 0),
|
||||
"total_tokens": getattr(response.usage, "total_tokens", 0),
|
||||
}
|
||||
usage_stats = TokenUsage(
|
||||
input_tokens=getattr(response.usage, "prompt_tokens", 0),
|
||||
output_tokens=getattr(response.usage, "completion_tokens", 0),
|
||||
)
|
||||
|
||||
latency = end_time - start_time
|
||||
|
||||
@@ -521,7 +523,7 @@ class WrappedEmbeddings:
|
||||
sanitize_openai_response(kwargs.get("input")),
|
||||
),
|
||||
"$ai_http_status": 200,
|
||||
"$ai_input_tokens": usage_stats.get("prompt_tokens", 0),
|
||||
"$ai_input_tokens": usage_stats.get("input_tokens", 0),
|
||||
"$ai_latency": latency,
|
||||
"$ai_trace_id": posthog_trace_id,
|
||||
"$ai_base_url": str(self._client.base_url),
|
||||
|
||||
@@ -14,7 +14,6 @@ from posthog.ai.types import (
|
||||
FormattedImageContent,
|
||||
FormattedMessage,
|
||||
FormattedTextContent,
|
||||
StreamingUsageStats,
|
||||
TokenUsage,
|
||||
)
|
||||
|
||||
@@ -256,9 +255,69 @@ def format_openai_streaming_content(
|
||||
return formatted
|
||||
|
||||
|
||||
def extract_openai_usage_from_response(response: Any) -> TokenUsage:
|
||||
"""
|
||||
Extract usage statistics from a full OpenAI response (non-streaming).
|
||||
Handles both Chat Completions and Responses API.
|
||||
|
||||
Args:
|
||||
response: The complete response from OpenAI API
|
||||
|
||||
Returns:
|
||||
TokenUsage with standardized usage statistics
|
||||
"""
|
||||
if not hasattr(response, "usage"):
|
||||
return TokenUsage(input_tokens=0, output_tokens=0)
|
||||
|
||||
cached_tokens = 0
|
||||
input_tokens = 0
|
||||
output_tokens = 0
|
||||
reasoning_tokens = 0
|
||||
|
||||
# Responses API format
|
||||
if hasattr(response.usage, "input_tokens"):
|
||||
input_tokens = response.usage.input_tokens
|
||||
if hasattr(response.usage, "output_tokens"):
|
||||
output_tokens = response.usage.output_tokens
|
||||
if hasattr(response.usage, "input_tokens_details") and hasattr(
|
||||
response.usage.input_tokens_details, "cached_tokens"
|
||||
):
|
||||
cached_tokens = response.usage.input_tokens_details.cached_tokens
|
||||
if hasattr(response.usage, "output_tokens_details") and hasattr(
|
||||
response.usage.output_tokens_details, "reasoning_tokens"
|
||||
):
|
||||
reasoning_tokens = response.usage.output_tokens_details.reasoning_tokens
|
||||
|
||||
# Chat Completions format
|
||||
if hasattr(response.usage, "prompt_tokens"):
|
||||
input_tokens = response.usage.prompt_tokens
|
||||
if hasattr(response.usage, "completion_tokens"):
|
||||
output_tokens = response.usage.completion_tokens
|
||||
if hasattr(response.usage, "prompt_tokens_details") and hasattr(
|
||||
response.usage.prompt_tokens_details, "cached_tokens"
|
||||
):
|
||||
cached_tokens = response.usage.prompt_tokens_details.cached_tokens
|
||||
if hasattr(response.usage, "completion_tokens_details") and hasattr(
|
||||
response.usage.completion_tokens_details, "reasoning_tokens"
|
||||
):
|
||||
reasoning_tokens = response.usage.completion_tokens_details.reasoning_tokens
|
||||
|
||||
result = TokenUsage(
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
)
|
||||
|
||||
if cached_tokens > 0:
|
||||
result["cache_read_input_tokens"] = cached_tokens
|
||||
if reasoning_tokens > 0:
|
||||
result["reasoning_tokens"] = reasoning_tokens
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def extract_openai_usage_from_chunk(
|
||||
chunk: Any, provider_type: str = "chat"
|
||||
) -> StreamingUsageStats:
|
||||
) -> TokenUsage:
|
||||
"""
|
||||
Extract usage statistics from an OpenAI streaming chunk.
|
||||
|
||||
@@ -272,16 +331,16 @@ def extract_openai_usage_from_chunk(
|
||||
Dictionary of usage statistics
|
||||
"""
|
||||
|
||||
usage: StreamingUsageStats = {}
|
||||
usage: TokenUsage = TokenUsage()
|
||||
|
||||
if provider_type == "chat":
|
||||
if not hasattr(chunk, "usage") or not chunk.usage:
|
||||
return usage
|
||||
|
||||
# Chat Completions API uses prompt_tokens and completion_tokens
|
||||
usage["prompt_tokens"] = getattr(chunk.usage, "prompt_tokens", 0)
|
||||
usage["completion_tokens"] = getattr(chunk.usage, "completion_tokens", 0)
|
||||
usage["total_tokens"] = getattr(chunk.usage, "total_tokens", 0)
|
||||
# Standardize to input_tokens and output_tokens
|
||||
usage["input_tokens"] = getattr(chunk.usage, "prompt_tokens", 0)
|
||||
usage["output_tokens"] = getattr(chunk.usage, "completion_tokens", 0)
|
||||
|
||||
# Handle cached tokens
|
||||
if hasattr(chunk.usage, "prompt_tokens_details") and hasattr(
|
||||
@@ -310,7 +369,6 @@ def extract_openai_usage_from_chunk(
|
||||
response_usage = chunk.response.usage
|
||||
usage["input_tokens"] = getattr(response_usage, "input_tokens", 0)
|
||||
usage["output_tokens"] = getattr(response_usage, "output_tokens", 0)
|
||||
usage["total_tokens"] = getattr(response_usage, "total_tokens", 0)
|
||||
|
||||
# Handle cached tokens
|
||||
if hasattr(response_usage, "input_tokens_details") and hasattr(
|
||||
@@ -535,37 +593,6 @@ def format_openai_streaming_output(
|
||||
]
|
||||
|
||||
|
||||
def standardize_openai_usage(
|
||||
usage: Dict[str, Any], api_type: str = "chat"
|
||||
) -> TokenUsage:
|
||||
"""
|
||||
Standardize OpenAI usage statistics to common TokenUsage format.
|
||||
|
||||
Args:
|
||||
usage: Raw usage statistics from OpenAI
|
||||
api_type: Either "chat" or "responses" to handle different field names
|
||||
|
||||
Returns:
|
||||
Standardized TokenUsage dict
|
||||
"""
|
||||
if api_type == "chat":
|
||||
# Chat API uses prompt_tokens/completion_tokens
|
||||
return TokenUsage(
|
||||
input_tokens=usage.get("prompt_tokens", 0),
|
||||
output_tokens=usage.get("completion_tokens", 0),
|
||||
cache_read_input_tokens=usage.get("cache_read_input_tokens"),
|
||||
reasoning_tokens=usage.get("reasoning_tokens"),
|
||||
)
|
||||
else: # responses API
|
||||
# Responses API uses input_tokens/output_tokens
|
||||
return TokenUsage(
|
||||
input_tokens=usage.get("input_tokens", 0),
|
||||
output_tokens=usage.get("output_tokens", 0),
|
||||
cache_read_input_tokens=usage.get("cache_read_input_tokens"),
|
||||
reasoning_tokens=usage.get("reasoning_tokens"),
|
||||
)
|
||||
|
||||
|
||||
def format_openai_streaming_input(
|
||||
kwargs: Dict[str, Any], api_type: str = "chat"
|
||||
) -> Any:
|
||||
@@ -579,7 +606,6 @@ def format_openai_streaming_input(
|
||||
Returns:
|
||||
Formatted input ready for PostHog tracking
|
||||
"""
|
||||
if api_type == "chat":
|
||||
return kwargs.get("messages")
|
||||
else: # responses API
|
||||
return kwargs.get("input")
|
||||
from posthog.ai.utils import merge_system_prompt
|
||||
|
||||
return merge_system_prompt(kwargs, "openai")
|
||||
|
||||
+1
-19
@@ -77,24 +77,6 @@ class ProviderResponse(TypedDict, total=False):
|
||||
error: Optional[str]
|
||||
|
||||
|
||||
class StreamingUsageStats(TypedDict, total=False):
|
||||
"""
|
||||
Usage statistics collected during streaming.
|
||||
|
||||
Different providers populate different fields during streaming.
|
||||
"""
|
||||
|
||||
input_tokens: int
|
||||
output_tokens: int
|
||||
cache_read_input_tokens: Optional[int]
|
||||
cache_creation_input_tokens: Optional[int]
|
||||
reasoning_tokens: Optional[int]
|
||||
# OpenAI-specific names
|
||||
prompt_tokens: Optional[int]
|
||||
completion_tokens: Optional[int]
|
||||
total_tokens: Optional[int]
|
||||
|
||||
|
||||
class StreamingContentBlock(TypedDict, total=False):
|
||||
"""
|
||||
Content block used during streaming to accumulate content.
|
||||
@@ -133,7 +115,7 @@ class StreamingEventData(TypedDict):
|
||||
kwargs: Dict[str, Any] # Original kwargs for tool extraction and special handling
|
||||
formatted_input: Any # Provider-formatted input ready for tracking
|
||||
formatted_output: Any # Provider-formatted output ready for tracking
|
||||
usage_stats: TokenUsage # Standardized token counts
|
||||
usage_stats: TokenUsage
|
||||
latency: float
|
||||
distinct_id: Optional[str]
|
||||
trace_id: Optional[str]
|
||||
|
||||
+94
-118
@@ -1,10 +1,9 @@
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Callable, Dict, Optional
|
||||
|
||||
from typing import Any, Callable, Dict, List, Optional, cast
|
||||
|
||||
from posthog.client import Client as PostHogClient
|
||||
from posthog.ai.types import StreamingEventData, StreamingUsageStats
|
||||
from posthog.ai.types import FormattedMessage, StreamingEventData, TokenUsage
|
||||
from posthog.ai.sanitization import (
|
||||
sanitize_openai,
|
||||
sanitize_anthropic,
|
||||
@@ -14,7 +13,7 @@ from posthog.ai.sanitization import (
|
||||
|
||||
|
||||
def merge_usage_stats(
|
||||
target: Dict[str, int], source: StreamingUsageStats, mode: str = "incremental"
|
||||
target: TokenUsage, source: TokenUsage, mode: str = "incremental"
|
||||
) -> None:
|
||||
"""
|
||||
Merge streaming usage statistics into target dict, handling None values.
|
||||
@@ -25,19 +24,49 @@ def merge_usage_stats(
|
||||
|
||||
Args:
|
||||
target: Dictionary to update with usage stats
|
||||
source: StreamingUsageStats that may contain None values
|
||||
source: TokenUsage that may contain None values
|
||||
mode: Either "incremental" or "cumulative"
|
||||
"""
|
||||
if mode == "incremental":
|
||||
# Add new values to existing totals
|
||||
for key, value in source.items():
|
||||
if value is not None and isinstance(value, int):
|
||||
target[key] = target.get(key, 0) + value
|
||||
source_input = source.get("input_tokens")
|
||||
if source_input is not None:
|
||||
current = target.get("input_tokens") or 0
|
||||
target["input_tokens"] = current + source_input
|
||||
|
||||
source_output = source.get("output_tokens")
|
||||
if source_output is not None:
|
||||
current = target.get("output_tokens") or 0
|
||||
target["output_tokens"] = current + source_output
|
||||
|
||||
source_cache_read = source.get("cache_read_input_tokens")
|
||||
if source_cache_read is not None:
|
||||
current = target.get("cache_read_input_tokens") or 0
|
||||
target["cache_read_input_tokens"] = current + source_cache_read
|
||||
|
||||
source_cache_creation = source.get("cache_creation_input_tokens")
|
||||
if source_cache_creation is not None:
|
||||
current = target.get("cache_creation_input_tokens") or 0
|
||||
target["cache_creation_input_tokens"] = current + source_cache_creation
|
||||
|
||||
source_reasoning = source.get("reasoning_tokens")
|
||||
if source_reasoning is not None:
|
||||
current = target.get("reasoning_tokens") or 0
|
||||
target["reasoning_tokens"] = current + source_reasoning
|
||||
elif mode == "cumulative":
|
||||
# Replace with latest values (already cumulative)
|
||||
for key, value in source.items():
|
||||
if value is not None and isinstance(value, int):
|
||||
target[key] = value
|
||||
if source.get("input_tokens") is not None:
|
||||
target["input_tokens"] = source["input_tokens"]
|
||||
if source.get("output_tokens") is not None:
|
||||
target["output_tokens"] = source["output_tokens"]
|
||||
if source.get("cache_read_input_tokens") is not None:
|
||||
target["cache_read_input_tokens"] = source["cache_read_input_tokens"]
|
||||
if source.get("cache_creation_input_tokens") is not None:
|
||||
target["cache_creation_input_tokens"] = source[
|
||||
"cache_creation_input_tokens"
|
||||
]
|
||||
if source.get("reasoning_tokens") is not None:
|
||||
target["reasoning_tokens"] = source["reasoning_tokens"]
|
||||
else:
|
||||
raise ValueError(f"Invalid mode: {mode}. Must be 'incremental' or 'cumulative'")
|
||||
|
||||
@@ -64,74 +93,31 @@ def get_model_params(kwargs: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return model_params
|
||||
|
||||
|
||||
def get_usage(response, provider: str) -> Dict[str, Any]:
|
||||
def get_usage(response, provider: str) -> TokenUsage:
|
||||
"""
|
||||
Extract usage statistics from response based on provider.
|
||||
Delegates to provider-specific converter functions.
|
||||
"""
|
||||
if provider == "anthropic":
|
||||
return {
|
||||
"input_tokens": response.usage.input_tokens,
|
||||
"output_tokens": response.usage.output_tokens,
|
||||
"cache_read_input_tokens": response.usage.cache_read_input_tokens,
|
||||
"cache_creation_input_tokens": response.usage.cache_creation_input_tokens,
|
||||
}
|
||||
from posthog.ai.anthropic.anthropic_converter import (
|
||||
extract_anthropic_usage_from_response,
|
||||
)
|
||||
|
||||
return extract_anthropic_usage_from_response(response)
|
||||
elif provider == "openai":
|
||||
cached_tokens = 0
|
||||
input_tokens = 0
|
||||
output_tokens = 0
|
||||
reasoning_tokens = 0
|
||||
from posthog.ai.openai.openai_converter import (
|
||||
extract_openai_usage_from_response,
|
||||
)
|
||||
|
||||
# responses api
|
||||
if hasattr(response.usage, "input_tokens"):
|
||||
input_tokens = response.usage.input_tokens
|
||||
if hasattr(response.usage, "output_tokens"):
|
||||
output_tokens = response.usage.output_tokens
|
||||
if hasattr(response.usage, "input_tokens_details") and hasattr(
|
||||
response.usage.input_tokens_details, "cached_tokens"
|
||||
):
|
||||
cached_tokens = response.usage.input_tokens_details.cached_tokens
|
||||
if hasattr(response.usage, "output_tokens_details") and hasattr(
|
||||
response.usage.output_tokens_details, "reasoning_tokens"
|
||||
):
|
||||
reasoning_tokens = response.usage.output_tokens_details.reasoning_tokens
|
||||
|
||||
# chat completions
|
||||
if hasattr(response.usage, "prompt_tokens"):
|
||||
input_tokens = response.usage.prompt_tokens
|
||||
if hasattr(response.usage, "completion_tokens"):
|
||||
output_tokens = response.usage.completion_tokens
|
||||
if hasattr(response.usage, "prompt_tokens_details") and hasattr(
|
||||
response.usage.prompt_tokens_details, "cached_tokens"
|
||||
):
|
||||
cached_tokens = response.usage.prompt_tokens_details.cached_tokens
|
||||
|
||||
return {
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
"cache_read_input_tokens": cached_tokens,
|
||||
"reasoning_tokens": reasoning_tokens,
|
||||
}
|
||||
return extract_openai_usage_from_response(response)
|
||||
elif provider == "gemini":
|
||||
input_tokens = 0
|
||||
output_tokens = 0
|
||||
from posthog.ai.gemini.gemini_converter import (
|
||||
extract_gemini_usage_from_response,
|
||||
)
|
||||
|
||||
if hasattr(response, "usage_metadata") and response.usage_metadata:
|
||||
input_tokens = getattr(response.usage_metadata, "prompt_token_count", 0)
|
||||
output_tokens = getattr(
|
||||
response.usage_metadata, "candidates_token_count", 0
|
||||
)
|
||||
return extract_gemini_usage_from_response(response)
|
||||
|
||||
return {
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"reasoning_tokens": 0,
|
||||
}
|
||||
return {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_creation_input_tokens": 0,
|
||||
"reasoning_tokens": 0,
|
||||
}
|
||||
return TokenUsage(input_tokens=0, output_tokens=0)
|
||||
|
||||
|
||||
def format_response(response, provider: str):
|
||||
@@ -169,9 +155,12 @@ def extract_available_tool_calls(provider: str, kwargs: Dict[str, Any]):
|
||||
from posthog.ai.openai.openai_converter import extract_openai_tools
|
||||
|
||||
return extract_openai_tools(kwargs)
|
||||
return None
|
||||
|
||||
|
||||
def merge_system_prompt(kwargs: Dict[str, Any], provider: str):
|
||||
def merge_system_prompt(
|
||||
kwargs: Dict[str, Any], provider: str
|
||||
) -> List[FormattedMessage]:
|
||||
"""
|
||||
Merge system prompts and format messages for the given provider.
|
||||
"""
|
||||
@@ -182,14 +171,15 @@ def merge_system_prompt(kwargs: Dict[str, Any], provider: str):
|
||||
system = kwargs.get("system")
|
||||
return format_anthropic_input(messages, system)
|
||||
elif provider == "gemini":
|
||||
from posthog.ai.gemini.gemini_converter import format_gemini_input
|
||||
from posthog.ai.gemini.gemini_converter import format_gemini_input_with_system
|
||||
|
||||
contents = kwargs.get("contents", [])
|
||||
return format_gemini_input(contents)
|
||||
config = kwargs.get("config")
|
||||
return format_gemini_input_with_system(contents, config)
|
||||
elif provider == "openai":
|
||||
# For OpenAI, handle both Chat Completions and Responses API
|
||||
from posthog.ai.openai.openai_converter import format_openai_input
|
||||
|
||||
# For OpenAI, handle both Chat Completions and Responses API
|
||||
messages_param = kwargs.get("messages")
|
||||
input_param = kwargs.get("input")
|
||||
|
||||
@@ -200,9 +190,11 @@ def merge_system_prompt(kwargs: Dict[str, Any], provider: str):
|
||||
if kwargs.get("system") is not None:
|
||||
has_system = any(msg.get("role") == "system" for msg in messages)
|
||||
if not has_system:
|
||||
messages = [
|
||||
{"role": "system", "content": kwargs.get("system")}
|
||||
] + messages
|
||||
system_msg = cast(
|
||||
FormattedMessage,
|
||||
{"role": "system", "content": kwargs.get("system")},
|
||||
)
|
||||
messages = [system_msg] + messages
|
||||
|
||||
# For Responses API, add instructions to the system prompt if provided
|
||||
if kwargs.get("instructions") is not None:
|
||||
@@ -220,9 +212,11 @@ def merge_system_prompt(kwargs: Dict[str, Any], provider: str):
|
||||
)
|
||||
else:
|
||||
# Create a new system message with instructions
|
||||
messages = [
|
||||
{"role": "system", "content": kwargs.get("instructions")}
|
||||
] + messages
|
||||
instruction_msg = cast(
|
||||
FormattedMessage,
|
||||
{"role": "system", "content": kwargs.get("instructions")},
|
||||
)
|
||||
messages = [instruction_msg] + messages
|
||||
|
||||
return messages
|
||||
|
||||
@@ -250,7 +244,7 @@ def call_llm_and_track_usage(
|
||||
response = None
|
||||
error = None
|
||||
http_status = 200
|
||||
usage: Dict[str, Any] = {}
|
||||
usage: TokenUsage = TokenUsage()
|
||||
error_params: Dict[str, Any] = {}
|
||||
|
||||
try:
|
||||
@@ -305,27 +299,17 @@ def call_llm_and_track_usage(
|
||||
if available_tool_calls:
|
||||
event_properties["$ai_tools"] = available_tool_calls
|
||||
|
||||
if (
|
||||
usage.get("cache_read_input_tokens") is not None
|
||||
and usage.get("cache_read_input_tokens", 0) > 0
|
||||
):
|
||||
event_properties["$ai_cache_read_input_tokens"] = usage.get(
|
||||
"cache_read_input_tokens", 0
|
||||
)
|
||||
cache_read = usage.get("cache_read_input_tokens")
|
||||
if cache_read is not None and cache_read > 0:
|
||||
event_properties["$ai_cache_read_input_tokens"] = cache_read
|
||||
|
||||
if (
|
||||
usage.get("cache_creation_input_tokens") is not None
|
||||
and usage.get("cache_creation_input_tokens", 0) > 0
|
||||
):
|
||||
event_properties["$ai_cache_creation_input_tokens"] = usage.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
cache_creation = usage.get("cache_creation_input_tokens")
|
||||
if cache_creation is not None and cache_creation > 0:
|
||||
event_properties["$ai_cache_creation_input_tokens"] = cache_creation
|
||||
|
||||
if (
|
||||
usage.get("reasoning_tokens") is not None
|
||||
and usage.get("reasoning_tokens", 0) > 0
|
||||
):
|
||||
event_properties["$ai_reasoning_tokens"] = usage.get("reasoning_tokens", 0)
|
||||
reasoning = usage.get("reasoning_tokens")
|
||||
if reasoning is not None and reasoning > 0:
|
||||
event_properties["$ai_reasoning_tokens"] = reasoning
|
||||
|
||||
if posthog_distinct_id is None:
|
||||
event_properties["$process_person_profile"] = False
|
||||
@@ -367,7 +351,7 @@ async def call_llm_and_track_usage_async(
|
||||
response = None
|
||||
error = None
|
||||
http_status = 200
|
||||
usage: Dict[str, Any] = {}
|
||||
usage: TokenUsage = TokenUsage()
|
||||
error_params: Dict[str, Any] = {}
|
||||
|
||||
try:
|
||||
@@ -422,21 +406,13 @@ async def call_llm_and_track_usage_async(
|
||||
if available_tool_calls:
|
||||
event_properties["$ai_tools"] = available_tool_calls
|
||||
|
||||
if (
|
||||
usage.get("cache_read_input_tokens") is not None
|
||||
and usage.get("cache_read_input_tokens", 0) > 0
|
||||
):
|
||||
event_properties["$ai_cache_read_input_tokens"] = usage.get(
|
||||
"cache_read_input_tokens", 0
|
||||
)
|
||||
cache_read = usage.get("cache_read_input_tokens")
|
||||
if cache_read is not None and cache_read > 0:
|
||||
event_properties["$ai_cache_read_input_tokens"] = cache_read
|
||||
|
||||
if (
|
||||
usage.get("cache_creation_input_tokens") is not None
|
||||
and usage.get("cache_creation_input_tokens", 0) > 0
|
||||
):
|
||||
event_properties["$ai_cache_creation_input_tokens"] = usage.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
cache_creation = usage.get("cache_creation_input_tokens")
|
||||
if cache_creation is not None and cache_creation > 0:
|
||||
event_properties["$ai_cache_creation_input_tokens"] = cache_creation
|
||||
|
||||
if posthog_distinct_id is None:
|
||||
event_properties["$process_person_profile"] = False
|
||||
|
||||
+31
-1
@@ -1,4 +1,4 @@
|
||||
from typing import TypedDict, Optional, Any, Dict, Union, Tuple, Type
|
||||
from typing import TypedDict, Optional, Any, Dict, List, Union, Tuple, Type
|
||||
from types import TracebackType
|
||||
from typing_extensions import NotRequired # For Python < 3.11 compatibility
|
||||
from datetime import datetime
|
||||
@@ -69,3 +69,33 @@ ExcInfo = Union[
|
||||
]
|
||||
|
||||
ExceptionArg = Union[BaseException, ExcInfo]
|
||||
|
||||
|
||||
# AI Event Types (literal strings to enforce valid event types)
|
||||
AI_EVENT_TYPE = Union[
|
||||
str, # Allow str for flexibility but document the expected values
|
||||
] # "$ai_generation", "$ai_trace", "$ai_span", "$ai_embedding", "$ai_metric", "$ai_feedback"
|
||||
|
||||
|
||||
class OptionalCaptureAIArgs(TypedDict):
|
||||
"""Optional arguments for the capture_ai method.
|
||||
|
||||
Args:
|
||||
distinct_id: Unique identifier for the person associated with this AI event. If not set, the context
|
||||
distinct_id is used, if available, otherwise a UUID is generated.
|
||||
properties: Dictionary of AI event properties to track. Must include required properties for the event type.
|
||||
blob_properties: List of property names that should be sent as blobs (e.g., '$ai_input', '$ai_output_choices').
|
||||
These properties will be extracted from `properties` and sent as multipart blobs.
|
||||
timestamp: When the event occurred (defaults to current time)
|
||||
uuid: Unique identifier for this specific event. If not provided, one is generated.
|
||||
groups: Group identifiers to associate with this event (format: {group_type: group_key})
|
||||
disable_geoip: Whether to disable GeoIP lookup for this event.
|
||||
"""
|
||||
|
||||
distinct_id: NotRequired[Optional[ID_TYPES]]
|
||||
properties: NotRequired[Optional[Dict[str, Any]]]
|
||||
blob_properties: NotRequired[Optional[List[str]]]
|
||||
timestamp: NotRequired[Optional[Union[datetime, str]]]
|
||||
uuid: NotRequired[Optional[str]]
|
||||
groups: NotRequired[Optional[Dict[str, str]]]
|
||||
disable_geoip: NotRequired[Optional[bool]]
|
||||
|
||||
+348
-5
@@ -1,16 +1,25 @@
|
||||
import atexit
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Dict, Optional, Union
|
||||
from io import BytesIO
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from typing_extensions import Unpack
|
||||
from uuid import uuid4
|
||||
|
||||
from dateutil.tz import tzutc
|
||||
from six import string_types
|
||||
|
||||
from posthog.args import OptionalCaptureArgs, OptionalSetArgs, ID_TYPES, ExceptionArg
|
||||
from posthog.args import (
|
||||
OptionalCaptureArgs,
|
||||
OptionalSetArgs,
|
||||
OptionalCaptureAIArgs,
|
||||
AI_EVENT_TYPE,
|
||||
ID_TYPES,
|
||||
ExceptionArg,
|
||||
)
|
||||
from posthog.consumer import Consumer
|
||||
from posthog.exception_capture import ExceptionCapture
|
||||
from posthog.exception_utils import (
|
||||
@@ -20,7 +29,11 @@ from posthog.exception_utils import (
|
||||
exception_is_already_captured,
|
||||
mark_exception_as_captured,
|
||||
)
|
||||
from posthog.feature_flags import InconclusiveMatchError, match_feature_flag_properties
|
||||
from posthog.feature_flags import (
|
||||
InconclusiveMatchError,
|
||||
RequiresServerEvaluation,
|
||||
match_feature_flag_properties,
|
||||
)
|
||||
from posthog.poller import Poller
|
||||
from posthog.request import (
|
||||
DEFAULT_HOST,
|
||||
@@ -99,6 +112,134 @@ def add_context_tags(properties):
|
||||
return properties
|
||||
|
||||
|
||||
def _generate_multipart_boundary() -> str:
|
||||
"""Generate a random boundary string for multipart requests."""
|
||||
return f"----WebKitFormBoundary{uuid4().hex[:16]}"
|
||||
|
||||
|
||||
def _encode_multipart_part(
|
||||
boundary: str,
|
||||
name: str,
|
||||
content: bytes,
|
||||
content_type: str,
|
||||
filename: Optional[str] = None,
|
||||
) -> bytes:
|
||||
"""Encode a single part of a multipart/form-data request."""
|
||||
part = f"--{boundary}\r\n".encode("utf-8")
|
||||
|
||||
if filename:
|
||||
part += f'Content-Disposition: form-data; name="{name}"; filename="{filename}"\r\n'.encode(
|
||||
"utf-8"
|
||||
)
|
||||
else:
|
||||
part += f'Content-Disposition: form-data; name="{name}"\r\n'.encode("utf-8")
|
||||
|
||||
part += f"Content-Type: {content_type}\r\n\r\n".encode("utf-8")
|
||||
part += content
|
||||
part += b"\r\n"
|
||||
|
||||
return part
|
||||
|
||||
|
||||
def _build_multipart_body(
|
||||
event_data: Dict[str, Any],
|
||||
properties: Optional[Dict[str, Any]],
|
||||
blob_properties: Optional[List[str]],
|
||||
) -> Tuple[bytes, str]:
|
||||
"""
|
||||
Build a multipart/form-data body for AI capture endpoint.
|
||||
|
||||
Args:
|
||||
event_data: The event data (uuid, event, distinct_id, timestamp)
|
||||
properties: Event properties (small properties)
|
||||
blob_properties: List of property names to send as blobs
|
||||
|
||||
Returns:
|
||||
Tuple of (body_bytes, content_type_header)
|
||||
"""
|
||||
boundary = _generate_multipart_boundary()
|
||||
body = BytesIO()
|
||||
|
||||
# Part 1: Event (required, must be first)
|
||||
event_json = json.dumps(event_data).encode("utf-8")
|
||||
body.write(
|
||||
_encode_multipart_part(boundary, "event", event_json, "application/json")
|
||||
)
|
||||
|
||||
# Separate properties into small properties and blob properties
|
||||
small_properties = {}
|
||||
blob_data = {}
|
||||
|
||||
if properties:
|
||||
blob_property_names = set(blob_properties or [])
|
||||
for key, value in properties.items():
|
||||
if key in blob_property_names:
|
||||
blob_data[key] = value
|
||||
else:
|
||||
small_properties[key] = value
|
||||
|
||||
# Part 2: Event properties (if there are any small properties)
|
||||
if small_properties:
|
||||
properties_json = json.dumps(small_properties).encode("utf-8")
|
||||
body.write(
|
||||
_encode_multipart_part(
|
||||
boundary, "event.properties", properties_json, "application/json"
|
||||
)
|
||||
)
|
||||
|
||||
# Part 3+: Blob parts
|
||||
for property_name, property_value in blob_data.items():
|
||||
# Serialize the blob data as JSON
|
||||
blob_json = json.dumps(property_value).encode("utf-8")
|
||||
blob_filename = f"blob_{uuid4().hex[:8]}"
|
||||
|
||||
body.write(
|
||||
_encode_multipart_part(
|
||||
boundary,
|
||||
f"event.properties.{property_name}",
|
||||
blob_json,
|
||||
"application/json",
|
||||
filename=blob_filename,
|
||||
)
|
||||
)
|
||||
|
||||
# Final boundary
|
||||
body.write(f"--{boundary}--\r\n".encode("utf-8"))
|
||||
|
||||
body_bytes = body.getvalue()
|
||||
content_type = f"multipart/form-data; boundary={boundary}"
|
||||
|
||||
return body_bytes, content_type
|
||||
|
||||
|
||||
def no_throw(default_return=None):
|
||||
"""
|
||||
Decorator to prevent raising exceptions from public API methods.
|
||||
Note that this doesn't prevent errors from propagating via `on_error`.
|
||||
Exceptions will still be raised if the debug flag is enabled.
|
||||
|
||||
Args:
|
||||
default_return: Value to return on exception (default: None)
|
||||
"""
|
||||
|
||||
def decorator(func):
|
||||
from functools import wraps
|
||||
|
||||
@wraps(func)
|
||||
def wrapper(self, *args, **kwargs):
|
||||
try:
|
||||
return func(self, *args, **kwargs)
|
||||
except Exception as e:
|
||||
if self.debug:
|
||||
raise e
|
||||
self.log.exception(f"Error in {func.__name__}: {e}")
|
||||
return default_return
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
class Client(object):
|
||||
"""
|
||||
This is the SDK reference for the PostHog Python SDK.
|
||||
@@ -481,6 +622,7 @@ class Client(object):
|
||||
|
||||
return normalize_flags_response(resp_data)
|
||||
|
||||
@no_throw()
|
||||
def capture(
|
||||
self, event: str, **kwargs: Unpack[OptionalCaptureArgs]
|
||||
) -> Optional[str]:
|
||||
@@ -657,6 +799,7 @@ class Client(object):
|
||||
f"Expected bool or dict."
|
||||
)
|
||||
|
||||
@no_throw()
|
||||
def set(self, **kwargs: Unpack[OptionalSetArgs]) -> Optional[str]:
|
||||
"""
|
||||
Set properties on a person profile.
|
||||
@@ -690,6 +833,8 @@ class Client(object):
|
||||
|
||||
Category:
|
||||
Identification
|
||||
|
||||
Note: This method will not raise exceptions. Errors are logged.
|
||||
"""
|
||||
distinct_id = kwargs.get("distinct_id", None)
|
||||
properties = kwargs.get("properties", None)
|
||||
@@ -716,6 +861,7 @@ class Client(object):
|
||||
|
||||
return self._enqueue(msg, disable_geoip)
|
||||
|
||||
@no_throw()
|
||||
def set_once(self, **kwargs: Unpack[OptionalSetArgs]) -> Optional[str]:
|
||||
"""
|
||||
Set properties on a person profile only if they haven't been set before.
|
||||
@@ -734,6 +880,8 @@ class Client(object):
|
||||
|
||||
Category:
|
||||
Identification
|
||||
|
||||
Note: This method will not raise exceptions. Errors are logged.
|
||||
"""
|
||||
distinct_id = kwargs.get("distinct_id", None)
|
||||
properties = kwargs.get("properties", None)
|
||||
@@ -759,6 +907,7 @@ class Client(object):
|
||||
|
||||
return self._enqueue(msg, disable_geoip)
|
||||
|
||||
@no_throw()
|
||||
def group_identify(
|
||||
self,
|
||||
group_type: str,
|
||||
@@ -791,6 +940,8 @@ class Client(object):
|
||||
|
||||
Category:
|
||||
Identification
|
||||
|
||||
Note: This method will not raise exceptions. Errors are logged.
|
||||
"""
|
||||
properties = properties or {}
|
||||
|
||||
@@ -815,6 +966,7 @@ class Client(object):
|
||||
|
||||
return self._enqueue(msg, disable_geoip)
|
||||
|
||||
@no_throw()
|
||||
def alias(
|
||||
self,
|
||||
previous_id: str,
|
||||
@@ -840,6 +992,8 @@ class Client(object):
|
||||
|
||||
Category:
|
||||
Identification
|
||||
|
||||
Note: This method will not raise exceptions. Errors are logged.
|
||||
"""
|
||||
(distinct_id, personless) = get_identity_state(distinct_id)
|
||||
|
||||
@@ -932,7 +1086,6 @@ class Client(object):
|
||||
"value"
|
||||
),
|
||||
"$exception_list": all_exceptions_with_trace_and_in_app,
|
||||
"$exception_personURL": f"{remove_trailing_slash(self.raw_host)}/project/{self.api_key}/person/{distinct_id}",
|
||||
**properties,
|
||||
}
|
||||
|
||||
@@ -961,6 +1114,196 @@ class Client(object):
|
||||
except Exception as e:
|
||||
self.log.exception(f"Failed to capture exception: {e}")
|
||||
|
||||
@no_throw()
|
||||
def capture_ai(
|
||||
self, event: AI_EVENT_TYPE, **kwargs: Unpack[OptionalCaptureAIArgs]
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Capture an AI event to the dedicated AI endpoint with support for large payloads.
|
||||
|
||||
This method sends AI events (like $ai_generation, $ai_trace, etc.) to PostHog's
|
||||
specialized AI endpoint (/i/v0/ai) which supports large payloads through multipart/form-data
|
||||
and blob storage in S3.
|
||||
|
||||
Args:
|
||||
event: The AI event type. Must be one of: "$ai_generation", "$ai_trace", "$ai_span",
|
||||
"$ai_embedding", "$ai_metric", "$ai_feedback"
|
||||
distinct_id: The distinct ID of the user.
|
||||
properties: A dictionary of AI event properties. Must include required properties based on event type:
|
||||
- All events: "$ai_model" (required)
|
||||
- $ai_generation: "$ai_provider", "$ai_trace_id" (required)
|
||||
- $ai_trace: "$ai_trace_id" (required)
|
||||
- $ai_span: "$ai_trace_id", "$ai_span_id" (required)
|
||||
- $ai_embedding: "$ai_provider", "$ai_trace_id" (required)
|
||||
blob_properties: List of property names to send as blobs (large data stored in S3).
|
||||
Common blob properties: "$ai_input", "$ai_output_choices", "$ai_input_state", "$ai_output_state"
|
||||
If not provided, defaults to common blob properties based on event type.
|
||||
timestamp: The timestamp of the event.
|
||||
uuid: A unique identifier for the event.
|
||||
groups: A dictionary of group information.
|
||||
disable_geoip: Whether to disable GeoIP for this event.
|
||||
|
||||
Examples:
|
||||
```python
|
||||
# $ai_generation event with blobs
|
||||
posthog.capture_ai(
|
||||
"$ai_generation",
|
||||
distinct_id="user_123",
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_provider": "openai",
|
||||
"$ai_trace_id": "trace_abc123",
|
||||
"$ai_input": {
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello!"}
|
||||
]
|
||||
},
|
||||
"$ai_output_choices": {
|
||||
"choices": [{"message": {"role": "assistant", "content": "Hi there!"}}]
|
||||
},
|
||||
"$ai_completion_tokens": 150,
|
||||
"$ai_prompt_tokens": 50
|
||||
},
|
||||
blob_properties=["$ai_input", "$ai_output_choices"]
|
||||
)
|
||||
```
|
||||
|
||||
```python
|
||||
# $ai_trace event
|
||||
posthog.capture_ai(
|
||||
"$ai_trace",
|
||||
distinct_id="user_123",
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_trace_id": "trace_abc123"
|
||||
}
|
||||
)
|
||||
```
|
||||
|
||||
```python
|
||||
# $ai_metric event
|
||||
posthog.capture_ai(
|
||||
"$ai_metric",
|
||||
distinct_id="user_123",
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_trace_id": "trace_abc123",
|
||||
"$ai_metric_name": "accuracy",
|
||||
"$ai_metric_value": "0.95"
|
||||
}
|
||||
)
|
||||
```
|
||||
|
||||
Category:
|
||||
AI Events
|
||||
|
||||
Note: This method sends events synchronously to the AI endpoint, bypassing the queue system.
|
||||
"""
|
||||
import requests
|
||||
|
||||
if self.disabled:
|
||||
return None
|
||||
|
||||
distinct_id = kwargs.get("distinct_id", None)
|
||||
properties = kwargs.get("properties", None)
|
||||
timestamp = kwargs.get("timestamp", None)
|
||||
uuid = kwargs.get("uuid", None)
|
||||
groups = kwargs.get("groups", None)
|
||||
blob_properties = kwargs.get("blob_properties", None)
|
||||
|
||||
# Default blob properties based on event type if not provided
|
||||
if blob_properties is None:
|
||||
default_blob_properties = {
|
||||
"$ai_generation": ["$ai_input", "$ai_output_choices"],
|
||||
"$ai_trace": ["$ai_input_state", "$ai_output_state"],
|
||||
"$ai_span": ["$ai_input_state", "$ai_output_state"],
|
||||
"$ai_embedding": ["$ai_input"],
|
||||
}
|
||||
blob_properties = default_blob_properties.get(event, [])
|
||||
|
||||
properties = properties or {}
|
||||
|
||||
# Get distinct_id
|
||||
(distinct_id, personless) = get_identity_state(distinct_id)
|
||||
|
||||
if personless and "$process_person_profile" not in properties:
|
||||
properties["$process_person_profile"] = False
|
||||
|
||||
# Prepare timestamp
|
||||
if timestamp is None:
|
||||
timestamp = datetime.now(tz=tzutc())
|
||||
|
||||
timestamp = guess_timezone(timestamp)
|
||||
timestamp_str = timestamp.isoformat()
|
||||
|
||||
# Generate UUID if not provided
|
||||
if uuid:
|
||||
uuid_str = stringify_id(uuid)
|
||||
else:
|
||||
uuid_str = stringify_id(uuid4())
|
||||
|
||||
# Add groups to properties
|
||||
if groups:
|
||||
properties["$groups"] = groups
|
||||
|
||||
# Prepare event data (the main event part)
|
||||
event_data = {
|
||||
"uuid": uuid_str,
|
||||
"event": event,
|
||||
"distinct_id": stringify_id(distinct_id),
|
||||
"timestamp": timestamp_str,
|
||||
}
|
||||
|
||||
# Build multipart body
|
||||
body_bytes, content_type = _build_multipart_body(
|
||||
event_data, properties, blob_properties
|
||||
)
|
||||
|
||||
# Send request to AI endpoint
|
||||
url = remove_trailing_slash(self.host) + "/i/v0/ai"
|
||||
|
||||
headers = {
|
||||
"Content-Type": content_type,
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"User-Agent": f"posthog-python/{VERSION}",
|
||||
}
|
||||
|
||||
self.log.debug(f"Sending AI event to {url}: {event} (uuid: {uuid_str})")
|
||||
|
||||
try:
|
||||
response = requests.post(
|
||||
url,
|
||||
data=body_bytes,
|
||||
headers=headers,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
self.log.debug(f"AI event captured successfully: {event}")
|
||||
return uuid_str
|
||||
else:
|
||||
error_message = f"AI capture failed with status {response.status_code}"
|
||||
try:
|
||||
error_detail = response.json()
|
||||
error_message = f"{error_message}: {error_detail}"
|
||||
except Exception:
|
||||
error_message = f"{error_message}: {response.text}"
|
||||
|
||||
self.log.error(error_message)
|
||||
|
||||
if self.debug:
|
||||
raise Exception(error_message)
|
||||
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
self.log.exception(f"Error sending AI event: {e}")
|
||||
|
||||
if self.debug:
|
||||
raise e
|
||||
|
||||
return None
|
||||
|
||||
def _enqueue(self, msg, disable_geoip):
|
||||
# type: (...) -> Optional[str]
|
||||
"""Push a new `msg` onto the queue, return `(success, msg)`"""
|
||||
@@ -1543,7 +1886,7 @@ class Client(object):
|
||||
self.log.debug(
|
||||
f"Successfully computed flag locally: {key} -> {response}"
|
||||
)
|
||||
except InconclusiveMatchError as e:
|
||||
except (RequiresServerEvaluation, InconclusiveMatchError) as e:
|
||||
self.log.debug(f"Failed to compute flag {key} locally: {e}")
|
||||
except Exception as e:
|
||||
self.log.exception(
|
||||
|
||||
+26
-10
@@ -22,6 +22,18 @@ class InconclusiveMatchError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class RequiresServerEvaluation(Exception):
|
||||
"""
|
||||
Raised when feature flag evaluation requires server-side data that is not
|
||||
available locally (e.g., static cohorts, experience continuity).
|
||||
|
||||
This error should propagate immediately to trigger API fallback, unlike
|
||||
InconclusiveMatchError which allows trying other conditions.
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
# This function takes a distinct_id and a feature flag key and returns a float between 0 and 1.
|
||||
# Given the same distinct_id and key, it'll always return the same float. These floats are
|
||||
# uniformly distributed between 0 and 1, so if we want to show this feature to 20% of traffic
|
||||
@@ -220,14 +232,7 @@ def match_feature_flag_properties(
|
||||
) or []
|
||||
valid_variant_keys = [variant["key"] for variant in flag_variants]
|
||||
|
||||
# Stable sort conditions with variant overrides to the top. This ensures that if overrides are present, they are
|
||||
# evaluated first, and the variant override is applied to the first matching condition.
|
||||
sorted_flag_conditions = sorted(
|
||||
flag_conditions,
|
||||
key=lambda condition: 0 if condition.get("variant") else 1,
|
||||
)
|
||||
|
||||
for condition in sorted_flag_conditions:
|
||||
for condition in flag_conditions:
|
||||
try:
|
||||
# if any one condition resolves to True, we can shortcircuit and return
|
||||
# the matching variant
|
||||
@@ -246,7 +251,12 @@ def match_feature_flag_properties(
|
||||
else:
|
||||
variant = get_matching_variant(flag, distinct_id)
|
||||
return variant or True
|
||||
except RequiresServerEvaluation:
|
||||
# Static cohort or other missing server-side data - must fallback to API
|
||||
raise
|
||||
except InconclusiveMatchError:
|
||||
# Evaluation error (bad regex, invalid date, missing property, etc.)
|
||||
# Track that we had an inconclusive match, but try other conditions
|
||||
is_inconclusive = True
|
||||
|
||||
if is_inconclusive:
|
||||
@@ -456,8 +466,8 @@ def match_cohort(
|
||||
# }
|
||||
cohort_id = str(property.get("value"))
|
||||
if cohort_id not in cohort_properties:
|
||||
raise InconclusiveMatchError(
|
||||
"can't match cohort without a given cohort property value"
|
||||
raise RequiresServerEvaluation(
|
||||
f"cohort {cohort_id} not found in local cohorts - likely a static cohort that requires server evaluation"
|
||||
)
|
||||
|
||||
property_group = cohort_properties[cohort_id]
|
||||
@@ -510,6 +520,9 @@ def match_property_group(
|
||||
# OR group
|
||||
if matches:
|
||||
return True
|
||||
except RequiresServerEvaluation:
|
||||
# Immediately propagate - this condition requires server-side data
|
||||
raise
|
||||
except InconclusiveMatchError as e:
|
||||
log.debug(f"Failed to compute property {prop} locally: {e}")
|
||||
error_matching_locally = True
|
||||
@@ -559,6 +572,9 @@ def match_property_group(
|
||||
return True
|
||||
if not matches and negation:
|
||||
return True
|
||||
except RequiresServerEvaluation:
|
||||
# Immediately propagate - this condition requires server-side data
|
||||
raise
|
||||
except InconclusiveMatchError as e:
|
||||
log.debug(f"Failed to compute property {prop} locally: {e}")
|
||||
error_matching_locally = True
|
||||
|
||||
@@ -1,10 +1,24 @@
|
||||
from typing import TYPE_CHECKING, cast
|
||||
from posthog import contexts, capture_exception
|
||||
from posthog import contexts
|
||||
from posthog.client import Client
|
||||
|
||||
try:
|
||||
from asgiref.sync import iscoroutinefunction, markcoroutinefunction
|
||||
except ImportError:
|
||||
# Fallback for older Django versions without asgiref
|
||||
import asyncio
|
||||
|
||||
iscoroutinefunction = asyncio.iscoroutinefunction
|
||||
|
||||
# No-op fallback for markcoroutinefunction
|
||||
# Older Django versions without asgiref typically don't support async middleware anyway
|
||||
def markcoroutinefunction(func):
|
||||
return func
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from django.http import HttpRequest, HttpResponse # noqa: F401
|
||||
from typing import Callable, Dict, Any, Optional # noqa: F401
|
||||
from typing import Callable, Dict, Any, Optional, Union, Awaitable # noqa: F401
|
||||
|
||||
|
||||
class PosthogContextMiddleware:
|
||||
@@ -31,11 +45,24 @@ class PosthogContextMiddleware:
|
||||
See the context documentation for more information. The extracted distinct ID and session ID, if found, are used to
|
||||
associate all events captured in the middleware context with the same distinct ID and session as currently active on the
|
||||
frontend. See the documentation for `set_context_session` and `identify_context` for more details.
|
||||
|
||||
This middleware is hybrid-capable: it supports both WSGI (sync) and ASGI (async) Django applications. The middleware
|
||||
detects at initialization whether the next middleware in the chain is async or sync, and adapts its behavior accordingly.
|
||||
This ensures compatibility with both pure sync and pure async middleware chains, as well as mixed chains in ASGI mode.
|
||||
"""
|
||||
|
||||
sync_capable = True
|
||||
async_capable = True
|
||||
|
||||
def __init__(self, get_response):
|
||||
# type: (Callable[[HttpRequest], HttpResponse]) -> None
|
||||
# type: (Union[Callable[[HttpRequest], HttpResponse], Callable[[HttpRequest], Awaitable[HttpResponse]]]) -> None
|
||||
self.get_response = get_response
|
||||
self._is_coroutine = iscoroutinefunction(get_response)
|
||||
|
||||
# Mark this instance as a coroutine function if get_response is async
|
||||
# This is required for Django to correctly detect async middleware
|
||||
if self._is_coroutine:
|
||||
markcoroutinefunction(self)
|
||||
|
||||
from django.conf import settings
|
||||
|
||||
@@ -158,24 +185,67 @@ class PosthogContextMiddleware:
|
||||
return user_id, email
|
||||
|
||||
def __call__(self, request):
|
||||
# type: (HttpRequest) -> HttpResponse
|
||||
# type: (HttpRequest) -> Union[HttpResponse, Awaitable[HttpResponse]]
|
||||
"""
|
||||
Unified entry point for both sync and async request handling.
|
||||
|
||||
When sync_capable and async_capable are both True, Django passes requests
|
||||
without conversion. This method detects the mode and routes accordingly.
|
||||
"""
|
||||
if self._is_coroutine:
|
||||
return self.__acall__(request)
|
||||
else:
|
||||
# Synchronous path
|
||||
if self.request_filter and not self.request_filter(request):
|
||||
return self.get_response(request)
|
||||
|
||||
with contexts.new_context(self.capture_exceptions, client=self.client):
|
||||
for k, v in self.extract_tags(request).items():
|
||||
contexts.tag(k, v)
|
||||
|
||||
return self.get_response(request)
|
||||
|
||||
async def __acall__(self, request):
|
||||
# type: (HttpRequest) -> Awaitable[HttpResponse]
|
||||
"""
|
||||
Asynchronous entry point for async request handling.
|
||||
|
||||
This method is called when the middleware chain is async.
|
||||
"""
|
||||
if self.request_filter and not self.request_filter(request):
|
||||
return self.get_response(request)
|
||||
return await self.get_response(request)
|
||||
|
||||
with contexts.new_context(self.capture_exceptions, client=self.client):
|
||||
for k, v in self.extract_tags(request).items():
|
||||
contexts.tag(k, v)
|
||||
|
||||
return self.get_response(request)
|
||||
return await self.get_response(request)
|
||||
|
||||
def process_exception(self, request, exception):
|
||||
# type: (HttpRequest, Exception) -> None
|
||||
"""
|
||||
Process exceptions from views and downstream middleware.
|
||||
|
||||
Django calls this WHILE still inside the context created by __call__,
|
||||
so request tags have already been extracted and set. This method just
|
||||
needs to capture the exception directly.
|
||||
|
||||
Django converts view exceptions into responses before they propagate through
|
||||
the middleware stack, so the context manager in __call__/__acall__ never sees them.
|
||||
|
||||
Note: Django's process_exception is always synchronous, even for async views.
|
||||
"""
|
||||
if self.request_filter and not self.request_filter(request):
|
||||
return
|
||||
|
||||
if not self.capture_exceptions:
|
||||
return
|
||||
|
||||
# Context and tags already set by __call__ or __acall__
|
||||
# Just capture the exception
|
||||
if self.client:
|
||||
self.client.capture_exception(exception)
|
||||
else:
|
||||
from posthog import capture_exception
|
||||
|
||||
capture_exception(exception)
|
||||
|
||||
@@ -12,8 +12,6 @@ try:
|
||||
except ImportError:
|
||||
ANTHROPIC_AVAILABLE = False
|
||||
|
||||
ANTHROPIC_API_KEY = os.getenv("ANTHROPIC_API_KEY")
|
||||
|
||||
# Skip all tests if Anthropic is not available
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not ANTHROPIC_AVAILABLE, reason="Anthropic package is not available"
|
||||
@@ -373,7 +371,6 @@ def test_privacy_mode_global(mock_client, mock_anthropic_response):
|
||||
assert props["$ai_output_choices"] is None
|
||||
|
||||
|
||||
@pytest.mark.skipif(not ANTHROPIC_API_KEY, reason="ANTHROPIC_API_KEY is not set")
|
||||
def test_basic_integration(mock_client):
|
||||
"""Test basic non-streaming integration."""
|
||||
|
||||
@@ -415,7 +412,6 @@ def test_basic_integration(mock_client):
|
||||
assert isinstance(props["$ai_latency"], float)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not ANTHROPIC_API_KEY, reason="ANTHROPIC_API_KEY is not set")
|
||||
async def test_basic_async_integration(mock_client):
|
||||
"""Test async non-streaming integration."""
|
||||
|
||||
@@ -459,7 +455,6 @@ async def test_basic_async_integration(mock_client):
|
||||
assert isinstance(props["$ai_latency"], float)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not ANTHROPIC_API_KEY, reason="ANTHROPIC_API_KEY is not set")
|
||||
async def test_async_streaming_system_prompt(mock_client):
|
||||
"""Test async streaming with system prompt."""
|
||||
|
||||
|
||||
@@ -31,6 +31,9 @@ def mock_gemini_response():
|
||||
mock_usage = MagicMock()
|
||||
mock_usage.prompt_token_count = 20
|
||||
mock_usage.candidates_token_count = 10
|
||||
# Ensure cache and reasoning tokens are not present (not MagicMock)
|
||||
mock_usage.cached_content_token_count = 0
|
||||
mock_usage.thoughts_token_count = 0
|
||||
mock_response.usage_metadata = mock_usage
|
||||
|
||||
mock_candidate = MagicMock()
|
||||
@@ -64,6 +67,8 @@ def mock_gemini_response_with_function_calls():
|
||||
mock_usage = MagicMock()
|
||||
mock_usage.prompt_token_count = 25
|
||||
mock_usage.candidates_token_count = 15
|
||||
mock_usage.cached_content_token_count = 0
|
||||
mock_usage.thoughts_token_count = 0
|
||||
mock_response.usage_metadata = mock_usage
|
||||
|
||||
# Mock function call
|
||||
@@ -110,6 +115,8 @@ def mock_gemini_response_function_calls_only():
|
||||
mock_usage = MagicMock()
|
||||
mock_usage.prompt_token_count = 30
|
||||
mock_usage.candidates_token_count = 12
|
||||
mock_usage.cached_content_token_count = 0
|
||||
mock_usage.thoughts_token_count = 0
|
||||
mock_response.usage_metadata = mock_usage
|
||||
|
||||
# Mock function call
|
||||
@@ -180,6 +187,8 @@ def test_new_client_streaming_with_generate_content_stream(
|
||||
mock_usage1 = MagicMock()
|
||||
mock_usage1.prompt_token_count = 10
|
||||
mock_usage1.candidates_token_count = 5
|
||||
mock_usage1.cached_content_token_count = 0
|
||||
mock_usage1.thoughts_token_count = 0
|
||||
mock_chunk1.usage_metadata = mock_usage1
|
||||
|
||||
mock_chunk2 = MagicMock()
|
||||
@@ -187,6 +196,8 @@ def test_new_client_streaming_with_generate_content_stream(
|
||||
mock_usage2 = MagicMock()
|
||||
mock_usage2.prompt_token_count = 10
|
||||
mock_usage2.candidates_token_count = 10
|
||||
mock_usage2.cached_content_token_count = 0
|
||||
mock_usage2.thoughts_token_count = 0
|
||||
mock_chunk2.usage_metadata = mock_usage2
|
||||
|
||||
yield mock_chunk1
|
||||
@@ -235,6 +246,8 @@ def test_new_client_streaming_with_tools(mock_client, mock_google_genai_client):
|
||||
mock_usage1 = MagicMock()
|
||||
mock_usage1.prompt_token_count = 15
|
||||
mock_usage1.candidates_token_count = 5
|
||||
mock_usage1.cached_content_token_count = 0
|
||||
mock_usage1.thoughts_token_count = 0
|
||||
mock_chunk1.usage_metadata = mock_usage1
|
||||
|
||||
mock_chunk2 = MagicMock()
|
||||
@@ -242,6 +255,8 @@ def test_new_client_streaming_with_tools(mock_client, mock_google_genai_client):
|
||||
mock_usage2 = MagicMock()
|
||||
mock_usage2.prompt_token_count = 15
|
||||
mock_usage2.candidates_token_count = 10
|
||||
mock_usage2.cached_content_token_count = 0
|
||||
mock_usage2.thoughts_token_count = 0
|
||||
mock_chunk2.usage_metadata = mock_usage2
|
||||
|
||||
yield mock_chunk1
|
||||
@@ -601,6 +616,8 @@ def test_tool_use_response(mock_client, mock_google_genai_client, mock_gemini_re
|
||||
|
||||
mock_config = MagicMock()
|
||||
mock_config.tools = [mock_tool]
|
||||
# Explicitly specify this config doesn't have system_instruction
|
||||
del mock_config.system_instruction
|
||||
|
||||
response = client.models.generate_content(
|
||||
model="gemini-2.5-flash",
|
||||
@@ -730,3 +747,93 @@ def test_function_calls_only_no_content(
|
||||
assert props["$ai_input_tokens"] == 30
|
||||
assert props["$ai_output_tokens"] == 12
|
||||
assert props["$ai_http_status"] == 200
|
||||
|
||||
|
||||
def test_cache_and_reasoning_tokens(mock_client, mock_google_genai_client):
|
||||
"""Test that cache and reasoning tokens are properly extracted"""
|
||||
# Create a mock response with cache and reasoning tokens
|
||||
mock_response = MagicMock()
|
||||
mock_response.text = "Test response with cache"
|
||||
|
||||
mock_usage = MagicMock()
|
||||
mock_usage.prompt_token_count = 100
|
||||
mock_usage.candidates_token_count = 50
|
||||
mock_usage.cached_content_token_count = 30 # Cache tokens
|
||||
mock_usage.thoughts_token_count = 10 # Reasoning tokens
|
||||
mock_response.usage_metadata = mock_usage
|
||||
|
||||
# Mock candidates
|
||||
mock_candidate = MagicMock()
|
||||
mock_candidate.text = "Test response with cache"
|
||||
mock_response.candidates = [mock_candidate]
|
||||
|
||||
mock_google_genai_client.models.generate_content.return_value = mock_response
|
||||
|
||||
client = Client(api_key="test-key", posthog_client=mock_client)
|
||||
|
||||
response = client.models.generate_content(
|
||||
model="gemini-2.5-pro",
|
||||
contents="Test with cache",
|
||||
posthog_distinct_id="test-id",
|
||||
)
|
||||
|
||||
assert response == mock_response
|
||||
assert mock_client.capture.call_count == 1
|
||||
|
||||
call_args = mock_client.capture.call_args[1]
|
||||
props = call_args["properties"]
|
||||
|
||||
# Check that all token types are present
|
||||
assert props["$ai_input_tokens"] == 100
|
||||
assert props["$ai_output_tokens"] == 50
|
||||
assert props["$ai_cache_read_input_tokens"] == 30
|
||||
assert props["$ai_reasoning_tokens"] == 10
|
||||
|
||||
|
||||
def test_streaming_cache_and_reasoning_tokens(mock_client, mock_google_genai_client):
|
||||
"""Test that cache and reasoning tokens are properly extracted in streaming"""
|
||||
# Create mock chunks with cache and reasoning tokens
|
||||
chunk1 = MagicMock()
|
||||
chunk1.text = "Hello "
|
||||
chunk1_usage = MagicMock()
|
||||
chunk1_usage.prompt_token_count = 100
|
||||
chunk1_usage.candidates_token_count = 5
|
||||
chunk1_usage.cached_content_token_count = 30 # Cache tokens
|
||||
chunk1_usage.thoughts_token_count = 0
|
||||
chunk1.usage_metadata = chunk1_usage
|
||||
|
||||
chunk2 = MagicMock()
|
||||
chunk2.text = "world!"
|
||||
chunk2_usage = MagicMock()
|
||||
chunk2_usage.prompt_token_count = 100
|
||||
chunk2_usage.candidates_token_count = 10
|
||||
chunk2_usage.cached_content_token_count = 30 # Same cache tokens
|
||||
chunk2_usage.thoughts_token_count = 5 # Reasoning tokens
|
||||
chunk2.usage_metadata = chunk2_usage
|
||||
|
||||
mock_stream = iter([chunk1, chunk2])
|
||||
mock_google_genai_client.models.generate_content_stream.return_value = mock_stream
|
||||
|
||||
client = Client(api_key="test-key", posthog_client=mock_client)
|
||||
|
||||
response = client.models.generate_content_stream(
|
||||
model="gemini-2.5-pro",
|
||||
contents="Test streaming with cache",
|
||||
posthog_distinct_id="test-id",
|
||||
)
|
||||
|
||||
# Consume the stream
|
||||
result = list(response)
|
||||
assert len(result) == 2
|
||||
|
||||
# Check PostHog capture was called
|
||||
assert mock_client.capture.call_count == 1
|
||||
|
||||
call_args = mock_client.capture.call_args[1]
|
||||
props = call_args["properties"]
|
||||
|
||||
# Check that all token types are present (should use final chunk's usage)
|
||||
assert props["$ai_input_tokens"] == 100
|
||||
assert props["$ai_output_tokens"] == 10
|
||||
assert props["$ai_cache_read_input_tokens"] == 30
|
||||
assert props["$ai_reasoning_tokens"] == 5
|
||||
|
||||
@@ -204,6 +204,7 @@ def test_basic_chat_chain(mock_client, stream):
|
||||
# Generation is second
|
||||
assert generation_args["event"] == "$ai_generation"
|
||||
assert "distinct_id" in generation_args
|
||||
assert generation_props["$ai_framework"] == "langchain"
|
||||
assert "$ai_model" in generation_props
|
||||
assert "$ai_provider" in generation_props
|
||||
assert generation_props["$ai_input"] == [
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import time
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -34,6 +34,7 @@ try:
|
||||
)
|
||||
|
||||
from posthog.ai.openai import OpenAI
|
||||
from posthog.ai.openai.openai_async import AsyncOpenAI
|
||||
|
||||
OPENAI_AVAILABLE = True
|
||||
except ImportError:
|
||||
@@ -218,6 +219,106 @@ def mock_openai_response_with_cached_tokens():
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def streaming_tool_call_chunks():
|
||||
return [
|
||||
ChatCompletionChunk(
|
||||
id="chunk1",
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
created=1234567890,
|
||||
choices=[
|
||||
ChoiceChunk(
|
||||
index=0,
|
||||
delta=ChoiceDelta(
|
||||
role="assistant",
|
||||
tool_calls=[
|
||||
ChoiceDeltaToolCall(
|
||||
index=0,
|
||||
id="call_abc123",
|
||||
type="function",
|
||||
function=ChoiceDeltaToolCallFunction(
|
||||
name="get_weather",
|
||||
arguments='{"location": "',
|
||||
),
|
||||
)
|
||||
],
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
),
|
||||
ChatCompletionChunk(
|
||||
id="chunk2",
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
created=1234567891,
|
||||
choices=[
|
||||
ChoiceChunk(
|
||||
index=0,
|
||||
delta=ChoiceDelta(
|
||||
tool_calls=[
|
||||
ChoiceDeltaToolCall(
|
||||
index=0,
|
||||
id="call_abc123",
|
||||
type="function",
|
||||
function=ChoiceDeltaToolCallFunction(
|
||||
arguments='San Francisco"',
|
||||
),
|
||||
)
|
||||
],
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
),
|
||||
ChatCompletionChunk(
|
||||
id="chunk3",
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
created=1234567892,
|
||||
choices=[
|
||||
ChoiceChunk(
|
||||
index=0,
|
||||
delta=ChoiceDelta(
|
||||
tool_calls=[
|
||||
ChoiceDeltaToolCall(
|
||||
index=0,
|
||||
id="call_abc123",
|
||||
type="function",
|
||||
function=ChoiceDeltaToolCallFunction(
|
||||
arguments=', "unit": "celsius"}',
|
||||
),
|
||||
)
|
||||
],
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
),
|
||||
ChatCompletionChunk(
|
||||
id="chunk4",
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
created=1234567893,
|
||||
choices=[
|
||||
ChoiceChunk(
|
||||
index=0,
|
||||
delta=ChoiceDelta(
|
||||
content="The weather in San Francisco is 15°C.",
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
usage=CompletionUsage(
|
||||
prompt_tokens=20,
|
||||
completion_tokens=15,
|
||||
total_tokens=35,
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_openai_response_with_tool_calls():
|
||||
return ChatCompletion(
|
||||
@@ -734,109 +835,11 @@ def test_responses_api_tool_calls(mock_client, mock_responses_api_with_tool_call
|
||||
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 = [
|
||||
ChatCompletionChunk(
|
||||
id="chunk1",
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
created=1234567890,
|
||||
choices=[
|
||||
ChoiceChunk(
|
||||
index=0,
|
||||
delta=ChoiceDelta(
|
||||
role="assistant",
|
||||
tool_calls=[
|
||||
ChoiceDeltaToolCall(
|
||||
index=0,
|
||||
id="call_abc123",
|
||||
type="function",
|
||||
function=ChoiceDeltaToolCallFunction(
|
||||
name="get_weather",
|
||||
arguments='{"location": "',
|
||||
),
|
||||
)
|
||||
],
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
),
|
||||
ChatCompletionChunk(
|
||||
id="chunk2",
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
created=1234567891,
|
||||
choices=[
|
||||
ChoiceChunk(
|
||||
index=0,
|
||||
delta=ChoiceDelta(
|
||||
tool_calls=[
|
||||
ChoiceDeltaToolCall(
|
||||
index=0,
|
||||
id="call_abc123",
|
||||
type="function",
|
||||
function=ChoiceDeltaToolCallFunction(
|
||||
arguments='San Francisco"',
|
||||
),
|
||||
)
|
||||
],
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
),
|
||||
ChatCompletionChunk(
|
||||
id="chunk3",
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
created=1234567892,
|
||||
choices=[
|
||||
ChoiceChunk(
|
||||
index=0,
|
||||
delta=ChoiceDelta(
|
||||
tool_calls=[
|
||||
ChoiceDeltaToolCall(
|
||||
index=0,
|
||||
id="call_abc123",
|
||||
type="function",
|
||||
function=ChoiceDeltaToolCallFunction(
|
||||
arguments=', "unit": "celsius"}',
|
||||
),
|
||||
)
|
||||
],
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
),
|
||||
ChatCompletionChunk(
|
||||
id="chunk4",
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
created=1234567893,
|
||||
choices=[
|
||||
ChoiceChunk(
|
||||
index=0,
|
||||
delta=ChoiceDelta(
|
||||
content="The weather in San Francisco is 15°C.",
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
usage=CompletionUsage(
|
||||
prompt_tokens=20,
|
||||
completion_tokens=15,
|
||||
total_tokens=35,
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
def test_streaming_with_tool_calls(mock_client, streaming_tool_call_chunks):
|
||||
# Mock the create method to return our chunks
|
||||
with patch("openai.resources.chat.completions.Completions.create") as mock_create:
|
||||
# Set up the mock to return our chunks when iterated
|
||||
mock_create.return_value = tool_call_chunks
|
||||
mock_create.return_value = streaming_tool_call_chunks
|
||||
|
||||
client = OpenAI(api_key="test-key", posthog_client=mock_client)
|
||||
|
||||
@@ -865,7 +868,7 @@ def test_streaming_with_tool_calls(mock_client):
|
||||
|
||||
# Verify the chunks were returned correctly
|
||||
assert len(chunks) == 4
|
||||
assert chunks == tool_call_chunks
|
||||
assert chunks == streaming_tool_call_chunks
|
||||
|
||||
# Verify the capture was called with the right arguments
|
||||
assert mock_client.capture.call_count == 1
|
||||
@@ -1112,6 +1115,169 @@ def test_responses_api_streaming_with_tokens(mock_client):
|
||||
assert isinstance(props["$ai_latency"], float)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_chat_streaming_with_tool_calls(
|
||||
mock_client, streaming_tool_call_chunks
|
||||
):
|
||||
captured_kwargs = {}
|
||||
|
||||
async def mock_create(self, **kwargs):
|
||||
captured_kwargs["kwargs"] = kwargs
|
||||
|
||||
async def chunk_iterable():
|
||||
for chunk in streaming_tool_call_chunks:
|
||||
yield chunk
|
||||
|
||||
return chunk_iterable()
|
||||
|
||||
with patch(
|
||||
"openai.resources.chat.completions.AsyncCompletions.create", new=mock_create
|
||||
):
|
||||
client = AsyncOpenAI(api_key="test-key", posthog_client=mock_client)
|
||||
|
||||
response_stream = await client.chat.completions.create(
|
||||
model="gpt-4",
|
||||
messages=[
|
||||
{"role": "user", "content": "What's the weather in San Francisco?"}
|
||||
],
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather",
|
||||
"parameters": {},
|
||||
},
|
||||
}
|
||||
],
|
||||
stream=True,
|
||||
posthog_distinct_id="test-id",
|
||||
)
|
||||
|
||||
chunks = []
|
||||
async for chunk in response_stream:
|
||||
chunks.append(chunk)
|
||||
|
||||
kwargs = captured_kwargs["kwargs"]
|
||||
assert kwargs["stream_options"]["include_usage"] is True
|
||||
|
||||
assert len(chunks) == len(streaming_tool_call_chunks)
|
||||
assert chunks == streaming_tool_call_chunks
|
||||
|
||||
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_tokens"] == 15
|
||||
assert props["$ai_input_tokens"] == 20
|
||||
assert isinstance(props["$ai_latency"], float)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_responses_streaming_with_tokens(mock_client):
|
||||
from openai.types.responses import ResponseUsage
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
chunks = []
|
||||
|
||||
chunk1 = MagicMock()
|
||||
chunk1.type = "response.text.delta"
|
||||
chunk1.text = "Test "
|
||||
chunks.append(chunk1)
|
||||
|
||||
chunk2 = MagicMock()
|
||||
chunk2.type = "response.text.delta"
|
||||
chunk2.text = "response"
|
||||
chunks.append(chunk2)
|
||||
|
||||
chunk3 = MagicMock()
|
||||
chunk3.type = "response.completed"
|
||||
chunk3.response = MagicMock()
|
||||
chunk3.response.usage = ResponseUsage(
|
||||
input_tokens=25,
|
||||
output_tokens=30,
|
||||
total_tokens=55,
|
||||
input_tokens_details={"prompt_tokens": 25, "cached_tokens": 0},
|
||||
output_tokens_details={"reasoning_tokens": 0},
|
||||
)
|
||||
chunk3.response.output = ["Test response"]
|
||||
chunks.append(chunk3)
|
||||
|
||||
captured_kwargs = {}
|
||||
|
||||
async def mock_create(self, **kwargs):
|
||||
captured_kwargs["kwargs"] = kwargs
|
||||
|
||||
async def chunk_iterable():
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
return chunk_iterable()
|
||||
|
||||
with patch("openai.resources.responses.AsyncResponses.create", new=mock_create):
|
||||
client = AsyncOpenAI(api_key="test-key", posthog_client=mock_client)
|
||||
|
||||
response_stream = await client.responses.create(
|
||||
model="gpt-4o-mini",
|
||||
input=[{"role": "user", "content": "Test message"}],
|
||||
stream=True,
|
||||
posthog_distinct_id="test-id",
|
||||
posthog_properties={"test": "streaming"},
|
||||
)
|
||||
|
||||
async for _ in response_stream:
|
||||
pass
|
||||
|
||||
kwargs = captured_kwargs["kwargs"]
|
||||
assert "stream_options" not in kwargs
|
||||
|
||||
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_tokens"] == 25
|
||||
assert props["$ai_output_tokens"] == 30
|
||||
assert props["test"] == "streaming"
|
||||
assert isinstance(props["$ai_latency"], float)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_embeddings_create(mock_client, mock_embedding_response):
|
||||
mock_create = AsyncMock(return_value=mock_embedding_response)
|
||||
|
||||
with patch("openai.resources.embeddings.AsyncEmbeddings.create", new=mock_create):
|
||||
client = AsyncOpenAI(api_key="test-key", posthog_client=mock_client)
|
||||
|
||||
response = await client.embeddings.create(
|
||||
model="text-embedding-3-small",
|
||||
input="Hello world",
|
||||
posthog_distinct_id="test-id",
|
||||
posthog_properties={"foo": "bar"},
|
||||
)
|
||||
|
||||
assert response == mock_embedding_response
|
||||
assert mock_create.await_count == 1
|
||||
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_embedding"
|
||||
assert props["$ai_provider"] == "openai"
|
||||
assert props["$ai_model"] == "text-embedding-3-small"
|
||||
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(
|
||||
|
||||
@@ -0,0 +1,354 @@
|
||||
"""
|
||||
Tests for system prompt capture across all LLM providers.
|
||||
|
||||
This test suite ensures that system prompts are correctly captured in analytics
|
||||
regardless of how they're passed to the providers:
|
||||
- As first message in messages/contents array (standard format)
|
||||
- As separate system parameter (Anthropic, OpenAI)
|
||||
- As instructions parameter (OpenAI Responses API)
|
||||
- As system_instruction parameter (Gemini)
|
||||
"""
|
||||
|
||||
import time
|
||||
import unittest
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
|
||||
class TestSystemPromptCapture(unittest.TestCase):
|
||||
"""Test system prompt capture for all providers."""
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.test_system_prompt = "You are a helpful AI assistant."
|
||||
self.test_user_message = "Hello, how are you?"
|
||||
self.test_response = "I'm doing well, thank you!"
|
||||
|
||||
# Create mock PostHog client
|
||||
self.client = MagicMock()
|
||||
self.client.privacy_mode = False
|
||||
|
||||
def _assert_system_prompt_captured(self, captured_input):
|
||||
"""Helper to assert system prompt is correctly captured."""
|
||||
self.assertEqual(
|
||||
len(captured_input), 2, "Should have 2 messages (system + user)"
|
||||
)
|
||||
self.assertEqual(
|
||||
captured_input[0]["role"], "system", "First message should be system"
|
||||
)
|
||||
self.assertEqual(
|
||||
captured_input[0]["content"],
|
||||
self.test_system_prompt,
|
||||
"System content should match",
|
||||
)
|
||||
self.assertEqual(
|
||||
captured_input[1]["role"], "user", "Second message should be user"
|
||||
)
|
||||
self.assertEqual(
|
||||
captured_input[1]["content"],
|
||||
self.test_user_message,
|
||||
"User content should match",
|
||||
)
|
||||
|
||||
# OpenAI Tests
|
||||
def test_openai_messages_array_system_prompt(self):
|
||||
"""Test OpenAI with system prompt in messages array."""
|
||||
try:
|
||||
from posthog.ai.openai import OpenAI
|
||||
from openai.types.chat import ChatCompletion, ChatCompletionMessage
|
||||
from openai.types.chat.chat_completion import Choice
|
||||
from openai.types.completion_usage import CompletionUsage
|
||||
except ImportError:
|
||||
self.skipTest("OpenAI package not available")
|
||||
|
||||
mock_response = ChatCompletion(
|
||||
id="test",
|
||||
model="gpt-4",
|
||||
object="chat.completion",
|
||||
created=int(time.time()),
|
||||
choices=[
|
||||
Choice(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=ChatCompletionMessage(
|
||||
content=self.test_response, role="assistant"
|
||||
),
|
||||
)
|
||||
],
|
||||
usage=CompletionUsage(
|
||||
completion_tokens=10, prompt_tokens=20, total_tokens=30
|
||||
),
|
||||
)
|
||||
|
||||
with patch(
|
||||
"openai.resources.chat.completions.Completions.create",
|
||||
return_value=mock_response,
|
||||
):
|
||||
client = OpenAI(posthog_client=self.client, api_key="test")
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": self.test_system_prompt},
|
||||
{"role": "user", "content": self.test_user_message},
|
||||
]
|
||||
|
||||
client.chat.completions.create(
|
||||
model="gpt-4", messages=messages, posthog_distinct_id="test-user"
|
||||
)
|
||||
|
||||
self.assertEqual(len(self.client.capture.call_args_list), 1)
|
||||
properties = self.client.capture.call_args_list[0][1]["properties"]
|
||||
self._assert_system_prompt_captured(properties["$ai_input"])
|
||||
|
||||
def test_openai_separate_system_parameter(self):
|
||||
"""Test OpenAI with system prompt as separate parameter."""
|
||||
try:
|
||||
from posthog.ai.openai import OpenAI
|
||||
from openai.types.chat import ChatCompletion, ChatCompletionMessage
|
||||
from openai.types.chat.chat_completion import Choice
|
||||
from openai.types.completion_usage import CompletionUsage
|
||||
except ImportError:
|
||||
self.skipTest("OpenAI package not available")
|
||||
|
||||
mock_response = ChatCompletion(
|
||||
id="test",
|
||||
model="gpt-4",
|
||||
object="chat.completion",
|
||||
created=int(time.time()),
|
||||
choices=[
|
||||
Choice(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=ChatCompletionMessage(
|
||||
content=self.test_response, role="assistant"
|
||||
),
|
||||
)
|
||||
],
|
||||
usage=CompletionUsage(
|
||||
completion_tokens=10, prompt_tokens=20, total_tokens=30
|
||||
),
|
||||
)
|
||||
|
||||
with patch(
|
||||
"openai.resources.chat.completions.Completions.create",
|
||||
return_value=mock_response,
|
||||
):
|
||||
client = OpenAI(posthog_client=self.client, api_key="test")
|
||||
|
||||
messages = [{"role": "user", "content": self.test_user_message}]
|
||||
|
||||
client.chat.completions.create(
|
||||
model="gpt-4",
|
||||
messages=messages,
|
||||
system=self.test_system_prompt,
|
||||
posthog_distinct_id="test-user",
|
||||
)
|
||||
|
||||
self.assertEqual(len(self.client.capture.call_args_list), 1)
|
||||
properties = self.client.capture.call_args_list[0][1]["properties"]
|
||||
self._assert_system_prompt_captured(properties["$ai_input"])
|
||||
|
||||
def test_openai_streaming_system_parameter(self):
|
||||
"""Test OpenAI streaming with system parameter."""
|
||||
try:
|
||||
from posthog.ai.openai import OpenAI
|
||||
from openai.types.chat.chat_completion_chunk import ChatCompletionChunk
|
||||
from openai.types.chat.chat_completion_chunk import Choice as ChoiceChunk
|
||||
from openai.types.chat.chat_completion_chunk import ChoiceDelta
|
||||
from openai.types.completion_usage import CompletionUsage
|
||||
except ImportError:
|
||||
self.skipTest("OpenAI package not available")
|
||||
|
||||
chunk1 = ChatCompletionChunk(
|
||||
id="test",
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
created=int(time.time()),
|
||||
choices=[
|
||||
ChoiceChunk(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=ChoiceDelta(content="Hello", role="assistant"),
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
chunk2 = ChatCompletionChunk(
|
||||
id="test",
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
created=int(time.time()),
|
||||
choices=[
|
||||
ChoiceChunk(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=ChoiceDelta(content=" there!", role=None),
|
||||
)
|
||||
],
|
||||
usage=CompletionUsage(
|
||||
completion_tokens=10, prompt_tokens=20, total_tokens=30
|
||||
),
|
||||
)
|
||||
|
||||
with patch(
|
||||
"openai.resources.chat.completions.Completions.create",
|
||||
return_value=[chunk1, chunk2],
|
||||
):
|
||||
client = OpenAI(posthog_client=self.client, api_key="test")
|
||||
|
||||
messages = [{"role": "user", "content": self.test_user_message}]
|
||||
|
||||
response_generator = client.chat.completions.create(
|
||||
model="gpt-4",
|
||||
messages=messages,
|
||||
system=self.test_system_prompt,
|
||||
stream=True,
|
||||
posthog_distinct_id="test-user",
|
||||
)
|
||||
|
||||
list(response_generator) # Consume generator
|
||||
|
||||
self.assertEqual(len(self.client.capture.call_args_list), 1)
|
||||
properties = self.client.capture.call_args_list[0][1]["properties"]
|
||||
self._assert_system_prompt_captured(properties["$ai_input"])
|
||||
|
||||
# Anthropic Tests
|
||||
def test_anthropic_messages_array_system_prompt(self):
|
||||
"""Test Anthropic with system prompt in messages array."""
|
||||
try:
|
||||
from posthog.ai.anthropic import Anthropic
|
||||
except ImportError:
|
||||
self.skipTest("Anthropic package not available")
|
||||
|
||||
with patch("anthropic.resources.messages.Messages.create") as mock_create:
|
||||
mock_response = MagicMock()
|
||||
mock_response.usage.input_tokens = 20
|
||||
mock_response.usage.output_tokens = 10
|
||||
mock_response.usage.cache_read_input_tokens = None
|
||||
mock_response.usage.cache_creation_input_tokens = None
|
||||
mock_create.return_value = mock_response
|
||||
|
||||
client = Anthropic(posthog_client=self.client, api_key="test")
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": self.test_system_prompt},
|
||||
{"role": "user", "content": self.test_user_message},
|
||||
]
|
||||
|
||||
client.messages.create(
|
||||
model="claude-3-5-sonnet-20241022",
|
||||
messages=messages,
|
||||
posthog_distinct_id="test-user",
|
||||
)
|
||||
|
||||
self.assertEqual(len(self.client.capture.call_args_list), 1)
|
||||
properties = self.client.capture.call_args_list[0][1]["properties"]
|
||||
self._assert_system_prompt_captured(properties["$ai_input"])
|
||||
|
||||
def test_anthropic_separate_system_parameter(self):
|
||||
"""Test Anthropic with system prompt as separate parameter."""
|
||||
try:
|
||||
from posthog.ai.anthropic import Anthropic
|
||||
except ImportError:
|
||||
self.skipTest("Anthropic package not available")
|
||||
|
||||
with patch("anthropic.resources.messages.Messages.create") as mock_create:
|
||||
mock_response = MagicMock()
|
||||
mock_response.usage.input_tokens = 20
|
||||
mock_response.usage.output_tokens = 10
|
||||
mock_response.usage.cache_read_input_tokens = None
|
||||
mock_response.usage.cache_creation_input_tokens = None
|
||||
mock_create.return_value = mock_response
|
||||
|
||||
client = Anthropic(posthog_client=self.client, api_key="test")
|
||||
|
||||
messages = [{"role": "user", "content": self.test_user_message}]
|
||||
|
||||
client.messages.create(
|
||||
model="claude-3-5-sonnet-20241022",
|
||||
messages=messages,
|
||||
system=self.test_system_prompt,
|
||||
posthog_distinct_id="test-user",
|
||||
)
|
||||
|
||||
self.assertEqual(len(self.client.capture.call_args_list), 1)
|
||||
properties = self.client.capture.call_args_list[0][1]["properties"]
|
||||
self._assert_system_prompt_captured(properties["$ai_input"])
|
||||
|
||||
# Gemini Tests
|
||||
def test_gemini_contents_array_system_prompt(self):
|
||||
"""Test Gemini with system prompt in contents array."""
|
||||
try:
|
||||
from posthog.ai.gemini import Client
|
||||
except ImportError:
|
||||
self.skipTest("Gemini package not available")
|
||||
|
||||
with patch("google.genai.Client") as mock_genai_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.candidates = [MagicMock()]
|
||||
mock_response.candidates[0].content.parts = [MagicMock()]
|
||||
mock_response.candidates[0].content.parts[0].text = self.test_response
|
||||
mock_response.usage_metadata.prompt_token_count = 20
|
||||
mock_response.usage_metadata.candidates_token_count = 10
|
||||
mock_response.usage_metadata.cached_content_token_count = None
|
||||
mock_response.usage_metadata.thoughts_token_count = None
|
||||
|
||||
mock_client_instance = MagicMock()
|
||||
mock_models_instance = MagicMock()
|
||||
mock_models_instance.generate_content.return_value = mock_response
|
||||
mock_client_instance.models = mock_models_instance
|
||||
mock_genai_class.return_value = mock_client_instance
|
||||
|
||||
client = Client(posthog_client=self.client, api_key="test")
|
||||
|
||||
contents = [
|
||||
{"role": "system", "content": self.test_system_prompt},
|
||||
{"role": "user", "content": self.test_user_message},
|
||||
]
|
||||
|
||||
client.models.generate_content(
|
||||
model="gemini-2.0-flash",
|
||||
contents=contents,
|
||||
posthog_distinct_id="test-user",
|
||||
)
|
||||
|
||||
self.assertEqual(len(self.client.capture.call_args_list), 1)
|
||||
properties = self.client.capture.call_args_list[0][1]["properties"]
|
||||
self._assert_system_prompt_captured(properties["$ai_input"])
|
||||
|
||||
def test_gemini_system_instruction_parameter(self):
|
||||
"""Test Gemini with system_instruction in config parameter."""
|
||||
try:
|
||||
from posthog.ai.gemini import Client
|
||||
except ImportError:
|
||||
self.skipTest("Gemini package not available")
|
||||
|
||||
with patch("google.genai.Client") as mock_genai_class:
|
||||
mock_response = MagicMock()
|
||||
mock_response.candidates = [MagicMock()]
|
||||
mock_response.candidates[0].content.parts = [MagicMock()]
|
||||
mock_response.candidates[0].content.parts[0].text = self.test_response
|
||||
mock_response.usage_metadata.prompt_token_count = 20
|
||||
mock_response.usage_metadata.candidates_token_count = 10
|
||||
mock_response.usage_metadata.cached_content_token_count = None
|
||||
mock_response.usage_metadata.thoughts_token_count = None
|
||||
|
||||
mock_client_instance = MagicMock()
|
||||
mock_models_instance = MagicMock()
|
||||
mock_models_instance.generate_content.return_value = mock_response
|
||||
mock_client_instance.models = mock_models_instance
|
||||
mock_genai_class.return_value = mock_client_instance
|
||||
|
||||
client = Client(posthog_client=self.client, api_key="test")
|
||||
|
||||
contents = [{"role": "user", "content": self.test_user_message}]
|
||||
config = {"system_instruction": self.test_system_prompt}
|
||||
|
||||
client.models.generate_content(
|
||||
model="gemini-2.0-flash",
|
||||
contents=contents,
|
||||
config=config,
|
||||
posthog_distinct_id="test-user",
|
||||
)
|
||||
|
||||
self.assertEqual(len(self.client.capture.call_args_list), 1)
|
||||
properties = self.client.capture.call_args_list[0][1]["properties"]
|
||||
self._assert_system_prompt_captured(properties["$ai_input"])
|
||||
@@ -4,7 +4,21 @@ from posthog.contexts import (
|
||||
get_context_distinct_id,
|
||||
)
|
||||
import unittest
|
||||
from unittest.mock import Mock
|
||||
from unittest.mock import Mock, patch
|
||||
import asyncio
|
||||
|
||||
# Configure Django settings before importing middleware
|
||||
import django
|
||||
from django.conf import settings
|
||||
|
||||
if not settings.configured:
|
||||
settings.configure(
|
||||
DEBUG=True,
|
||||
SECRET_KEY="test-secret-key",
|
||||
INSTALLED_APPS=[],
|
||||
MIDDLEWARE=[],
|
||||
)
|
||||
django.setup()
|
||||
|
||||
from posthog.integrations.django import PosthogContextMiddleware
|
||||
|
||||
@@ -38,14 +52,33 @@ class TestPosthogContextMiddleware(unittest.TestCase):
|
||||
request_filter=None,
|
||||
tag_map=None,
|
||||
capture_exceptions=True,
|
||||
get_response=None,
|
||||
):
|
||||
"""Helper to create middleware instance without calling __init__"""
|
||||
middleware = PosthogContextMiddleware.__new__(PosthogContextMiddleware)
|
||||
middleware.get_response = Mock()
|
||||
middleware.extra_tags = extra_tags
|
||||
middleware.request_filter = request_filter
|
||||
middleware.tag_map = tag_map
|
||||
middleware.capture_exceptions = capture_exceptions
|
||||
"""Helper to create middleware instance with mock Django settings"""
|
||||
if get_response is None:
|
||||
get_response = Mock()
|
||||
|
||||
with patch("django.conf.settings") as mock_settings:
|
||||
# Configure mock settings
|
||||
mock_settings.POSTHOG_MW_EXTRA_TAGS = extra_tags
|
||||
mock_settings.POSTHOG_MW_REQUEST_FILTER = request_filter
|
||||
mock_settings.POSTHOG_MW_TAG_MAP = tag_map
|
||||
mock_settings.POSTHOG_MW_CAPTURE_EXCEPTIONS = capture_exceptions
|
||||
mock_settings.POSTHOG_MW_CLIENT = None
|
||||
|
||||
# Make hasattr work correctly
|
||||
def mock_hasattr(obj, name):
|
||||
return name in [
|
||||
"POSTHOG_MW_EXTRA_TAGS",
|
||||
"POSTHOG_MW_REQUEST_FILTER",
|
||||
"POSTHOG_MW_TAG_MAP",
|
||||
"POSTHOG_MW_CAPTURE_EXCEPTIONS",
|
||||
"POSTHOG_MW_CLIENT",
|
||||
]
|
||||
|
||||
with patch("builtins.hasattr", side_effect=mock_hasattr):
|
||||
middleware = PosthogContextMiddleware(get_response)
|
||||
|
||||
return middleware
|
||||
|
||||
def test_extract_tags_basic(self):
|
||||
@@ -168,6 +201,348 @@ class TestPosthogContextMiddleware(unittest.TestCase):
|
||||
|
||||
self.assertEqual(tags["$request_method"], "PATCH")
|
||||
|
||||
def test_process_exception_called_during_view_exception(self):
|
||||
"""
|
||||
Unit test verifying process_exception captures exceptions per Django's contract.
|
||||
|
||||
Since this is a library test (no Django runtime), we simulate how Django
|
||||
would invoke our middleware in production:
|
||||
1. Middleware.__call__ creates context with request tags
|
||||
2. View raises exception inside get_response
|
||||
3. Django's BaseHandler catches it, calls process_exception, returns error response
|
||||
4. Exception never propagates to middleware's context manager
|
||||
|
||||
We manually call process_exception to simulate Django's behavior - this is
|
||||
the only way to test the hook without a full Django integration test.
|
||||
"""
|
||||
mock_client = Mock()
|
||||
view_exception = ValueError("View raised this error")
|
||||
error_response = Mock(status_code=500)
|
||||
|
||||
def mock_get_response(request):
|
||||
# Simulate Django's exception handling: catches view exception,
|
||||
# calls process_exception hook if it exists, returns error response
|
||||
if hasattr(middleware, "process_exception"):
|
||||
middleware.process_exception(request, view_exception)
|
||||
return error_response
|
||||
|
||||
middleware = self.create_middleware(get_response=mock_get_response)
|
||||
middleware.client = mock_client
|
||||
|
||||
request = MockRequest(
|
||||
headers={"X-POSTHOG-DISTINCT-ID": "test-user"},
|
||||
method="POST",
|
||||
path="/api/endpoint",
|
||||
)
|
||||
response = middleware(request)
|
||||
|
||||
self.assertEqual(response.status_code, 500)
|
||||
mock_client.capture_exception.assert_called_once_with(view_exception)
|
||||
|
||||
def test_process_exception_respects_capture_exceptions_false(self):
|
||||
"""Verify process_exception respects capture_exceptions=False setting"""
|
||||
mock_client = Mock()
|
||||
view_exception = ValueError("Should not be captured")
|
||||
|
||||
def mock_get_response(request):
|
||||
if hasattr(middleware, "process_exception"):
|
||||
middleware.process_exception(request, view_exception)
|
||||
return Mock(status_code=500)
|
||||
|
||||
middleware = self.create_middleware(
|
||||
capture_exceptions=False, get_response=mock_get_response
|
||||
)
|
||||
middleware.client = mock_client
|
||||
|
||||
request = MockRequest()
|
||||
middleware(request)
|
||||
|
||||
mock_client.capture_exception.assert_not_called()
|
||||
|
||||
def test_process_exception_respects_request_filter(self):
|
||||
"""Verify process_exception respects request_filter setting"""
|
||||
mock_client = Mock()
|
||||
view_exception = ValueError("Should be filtered")
|
||||
|
||||
def mock_get_response(request):
|
||||
if hasattr(middleware, "process_exception"):
|
||||
middleware.process_exception(request, view_exception)
|
||||
return Mock(status_code=500)
|
||||
|
||||
middleware = self.create_middleware(
|
||||
request_filter=lambda req: False,
|
||||
capture_exceptions=True,
|
||||
get_response=mock_get_response,
|
||||
)
|
||||
middleware.client = mock_client
|
||||
|
||||
request = MockRequest()
|
||||
middleware(request)
|
||||
|
||||
mock_client.capture_exception.assert_not_called()
|
||||
|
||||
|
||||
class TestPosthogContextMiddlewareSync(unittest.TestCase):
|
||||
"""Test synchronous middleware behavior"""
|
||||
|
||||
def test_sync_middleware_call(self):
|
||||
"""Test that sync middleware correctly processes requests"""
|
||||
mock_response = Mock()
|
||||
get_response = Mock(return_value=mock_response)
|
||||
|
||||
# Create middleware with sync get_response
|
||||
middleware = PosthogContextMiddleware(get_response)
|
||||
|
||||
# Verify sync mode detected
|
||||
self.assertFalse(middleware._is_coroutine)
|
||||
|
||||
request = MockRequest(
|
||||
headers={"X-POSTHOG-SESSION-ID": "test-session"},
|
||||
method="GET",
|
||||
path="/test",
|
||||
)
|
||||
|
||||
with new_context():
|
||||
response = middleware(request)
|
||||
|
||||
# Verify response returned
|
||||
self.assertEqual(response, mock_response)
|
||||
get_response.assert_called_once_with(request)
|
||||
|
||||
def test_sync_middleware_with_filter(self):
|
||||
"""Test sync middleware respects request filter"""
|
||||
mock_response = Mock()
|
||||
get_response = Mock(return_value=mock_response)
|
||||
|
||||
# Create middleware with request filter that filters all requests
|
||||
request_filter = lambda req: False
|
||||
middleware = PosthogContextMiddleware.__new__(PosthogContextMiddleware)
|
||||
middleware.get_response = get_response
|
||||
middleware._is_coroutine = False
|
||||
middleware.request_filter = request_filter
|
||||
middleware.capture_exceptions = True
|
||||
middleware.client = None
|
||||
|
||||
request = MockRequest()
|
||||
|
||||
# Should skip context creation and return response directly
|
||||
response = middleware(request)
|
||||
self.assertEqual(response, mock_response)
|
||||
get_response.assert_called_once_with(request)
|
||||
|
||||
def test_view_exceptions_only_captured_via_process_exception(self):
|
||||
"""
|
||||
Demonstrates that process_exception is required to capture view exceptions.
|
||||
|
||||
In production Django, view exceptions don't propagate to middleware's context
|
||||
manager because Django's BaseHandler catches them first and converts them to
|
||||
error responses. Django provides the exception via process_exception hook instead.
|
||||
|
||||
This unit test proves:
|
||||
1. Context manager in __call__ never sees view exceptions (Django intercepts)
|
||||
2. Only process_exception can capture them
|
||||
3. Without process_exception, exceptions are silently lost (v6.7.5 regression)
|
||||
|
||||
We manually call process_exception to verify the hook works - in production,
|
||||
Django's BaseHandler would call it when a view raises.
|
||||
"""
|
||||
mock_client = Mock()
|
||||
get_response = Mock(return_value=Mock(status_code=500))
|
||||
|
||||
middleware = PosthogContextMiddleware(get_response)
|
||||
middleware.client = mock_client
|
||||
|
||||
def get_response_simulating_django(request):
|
||||
# Simulates Django behavior: view exception converted to error response,
|
||||
# never propagates to middleware's context manager
|
||||
return Mock(status_code=500)
|
||||
|
||||
middleware._sync_get_response = get_response_simulating_django
|
||||
|
||||
request = MockRequest()
|
||||
|
||||
response = middleware(request)
|
||||
self.assertEqual(response.status_code, 500)
|
||||
|
||||
# Context manager didn't capture anything - exception was intercepted by Django
|
||||
mock_client.capture_exception.assert_not_called()
|
||||
|
||||
# Verify process_exception hook exists and captures exceptions when called
|
||||
if hasattr(middleware, "process_exception"):
|
||||
exception = ValueError("View error")
|
||||
middleware.process_exception(request, exception)
|
||||
mock_client.capture_exception.assert_called_once_with(exception)
|
||||
else:
|
||||
self.fail(
|
||||
"process_exception missing - view exceptions will not be captured!"
|
||||
)
|
||||
|
||||
|
||||
class TestPosthogContextMiddlewareAsync(unittest.TestCase):
|
||||
"""Test asynchronous middleware behavior"""
|
||||
|
||||
def test_async_middleware_detection(self):
|
||||
"""Test that async get_response is correctly detected"""
|
||||
|
||||
async def async_get_response(request):
|
||||
return Mock()
|
||||
|
||||
middleware = PosthogContextMiddleware(async_get_response)
|
||||
|
||||
# Verify async mode detected
|
||||
self.assertTrue(middleware._is_coroutine)
|
||||
|
||||
def test_async_middleware_call(self):
|
||||
"""Test that async middleware correctly processes requests"""
|
||||
|
||||
async def run_test():
|
||||
mock_response = Mock()
|
||||
|
||||
async def async_get_response(request):
|
||||
return mock_response
|
||||
|
||||
middleware = PosthogContextMiddleware(async_get_response)
|
||||
|
||||
request = MockRequest(
|
||||
headers={"X-POSTHOG-SESSION-ID": "async-session"},
|
||||
method="POST",
|
||||
path="/async-test",
|
||||
)
|
||||
|
||||
with new_context():
|
||||
# Call should return the coroutine from __acall__
|
||||
result = middleware(request)
|
||||
|
||||
# Verify it's a coroutine
|
||||
self.assertTrue(asyncio.iscoroutine(result))
|
||||
|
||||
# Await the result
|
||||
response = await result
|
||||
self.assertEqual(response, mock_response)
|
||||
|
||||
asyncio.run(run_test())
|
||||
|
||||
def test_async_middleware_with_filter(self):
|
||||
"""Test async middleware respects request filter"""
|
||||
|
||||
async def run_test():
|
||||
mock_response = Mock()
|
||||
|
||||
async def async_get_response(request):
|
||||
return mock_response
|
||||
|
||||
# Properly initialize middleware
|
||||
middleware = PosthogContextMiddleware(async_get_response)
|
||||
# Override request filter after initialization
|
||||
middleware.request_filter = lambda req: False
|
||||
|
||||
request = MockRequest()
|
||||
|
||||
# Should skip context creation and return response directly
|
||||
result = middleware(request)
|
||||
response = await result
|
||||
self.assertEqual(response, mock_response)
|
||||
|
||||
asyncio.run(run_test())
|
||||
|
||||
def test_async_middleware_context_propagation(self):
|
||||
"""Test that async middleware properly propagates context"""
|
||||
|
||||
async def run_test():
|
||||
mock_response = Mock()
|
||||
|
||||
async def async_get_response(request):
|
||||
# Verify context is available during async processing
|
||||
session_id = get_context_session_id()
|
||||
self.assertEqual(session_id, "async-session-123")
|
||||
return mock_response
|
||||
|
||||
middleware = PosthogContextMiddleware(async_get_response)
|
||||
|
||||
request = MockRequest(
|
||||
headers={"X-POSTHOG-SESSION-ID": "async-session-123"},
|
||||
method="GET",
|
||||
)
|
||||
|
||||
with new_context():
|
||||
result = middleware(request)
|
||||
await result
|
||||
|
||||
asyncio.run(run_test())
|
||||
|
||||
def test_async_middleware_exception_capture(self):
|
||||
"""Test that async middleware captures exceptions during request processing"""
|
||||
|
||||
async def run_test():
|
||||
mock_client = Mock()
|
||||
|
||||
# Make async_get_response raise an exception
|
||||
async def raise_exception(request):
|
||||
raise ValueError("Async test exception")
|
||||
|
||||
# Properly initialize middleware
|
||||
middleware = PosthogContextMiddleware(raise_exception)
|
||||
middleware.client = mock_client # Override with mock client
|
||||
|
||||
request = MockRequest()
|
||||
|
||||
# Should capture exception and re-raise
|
||||
with self.assertRaises(ValueError):
|
||||
result = middleware(request)
|
||||
await result
|
||||
|
||||
# Verify exception was captured by middleware
|
||||
mock_client.capture_exception.assert_called_once()
|
||||
captured_exception = mock_client.capture_exception.call_args[0][0]
|
||||
self.assertIsInstance(captured_exception, ValueError)
|
||||
self.assertEqual(str(captured_exception), "Async test exception")
|
||||
|
||||
asyncio.run(run_test())
|
||||
|
||||
|
||||
class TestPosthogContextMiddlewareHybrid(unittest.TestCase):
|
||||
"""Test hybrid middleware behavior with mixed sync/async chains"""
|
||||
|
||||
def test_hybrid_flags_set(self):
|
||||
"""Test that both capability flags are set"""
|
||||
self.assertTrue(PosthogContextMiddleware.sync_capable)
|
||||
self.assertTrue(PosthogContextMiddleware.async_capable)
|
||||
|
||||
def test_sync_to_async_routing(self):
|
||||
"""Test that __call__ routes to __acall__ when async"""
|
||||
|
||||
async def run_test():
|
||||
async def async_get_response(request):
|
||||
return Mock()
|
||||
|
||||
middleware = PosthogContextMiddleware(async_get_response)
|
||||
|
||||
# Verify routing happens
|
||||
request = MockRequest()
|
||||
result = middleware(request)
|
||||
|
||||
# Should be a coroutine from __acall__
|
||||
self.assertTrue(asyncio.iscoroutine(result))
|
||||
await result # Clean up
|
||||
|
||||
asyncio.run(run_test())
|
||||
|
||||
def test_sync_path_direct_return(self):
|
||||
"""Test that sync path returns directly without coroutine"""
|
||||
mock_response = Mock()
|
||||
|
||||
def sync_get_response(request):
|
||||
return mock_response
|
||||
|
||||
middleware = PosthogContextMiddleware(sync_get_response)
|
||||
|
||||
request = MockRequest()
|
||||
result = middleware(request)
|
||||
|
||||
# Should NOT be a coroutine
|
||||
self.assertFalse(asyncio.iscoroutine(result))
|
||||
self.assertEqual(result, mock_response)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -2423,3 +2423,46 @@ class TestClient(unittest.TestCase):
|
||||
batch_data = mock_post.call_args[1]["batch"]
|
||||
msg = batch_data[0]
|
||||
self.assertEqual(msg["properties"]["$context_tags"], ["random_tag"])
|
||||
|
||||
@mock.patch(
|
||||
"posthog.client.Client._enqueue", side_effect=Exception("Unexpected error")
|
||||
)
|
||||
def test_methods_handle_exceptions(self, mock_enqueue):
|
||||
"""Test that all decorated methods handle exceptions gracefully."""
|
||||
client = Client("test-key")
|
||||
|
||||
test_cases = [
|
||||
("capture", ["test_event"], {}),
|
||||
("set", [], {"distinct_id": "some-id", "properties": {"a": "b"}}),
|
||||
("set_once", [], {"distinct_id": "some-id", "properties": {"a": "b"}}),
|
||||
("group_identify", ["group-type", "group-key"], {}),
|
||||
("alias", ["some-id", "new-id"], {}),
|
||||
]
|
||||
|
||||
for method_name, args, kwargs in test_cases:
|
||||
with self.subTest(method=method_name):
|
||||
method = getattr(client, method_name)
|
||||
result = method(*args, **kwargs)
|
||||
self.assertEqual(result, None)
|
||||
|
||||
@mock.patch(
|
||||
"posthog.client.Client._enqueue", side_effect=Exception("Expected error")
|
||||
)
|
||||
def test_debug_flag_re_raises_exceptions(self, mock_enqueue):
|
||||
"""Test that methods re-raise exceptions when debug=True."""
|
||||
client = Client("test-key", debug=True)
|
||||
|
||||
test_cases = [
|
||||
("capture", ["test_event"], {}),
|
||||
("set", [], {"distinct_id": "some-id", "properties": {"a": "b"}}),
|
||||
("set_once", [], {"distinct_id": "some-id", "properties": {"a": "b"}}),
|
||||
("group_identify", ["group-type", "group-key"], {}),
|
||||
("alias", ["some-id", "new-id"], {}),
|
||||
]
|
||||
|
||||
for method_name, args, kwargs in test_cases:
|
||||
with self.subTest(method=method_name):
|
||||
method = getattr(client, method_name)
|
||||
with self.assertRaises(Exception) as cm:
|
||||
method(*args, **kwargs)
|
||||
self.assertEqual(str(cm.exception), "Expected error")
|
||||
|
||||
@@ -2804,73 +2804,61 @@ class TestLocalEvaluation(unittest.TestCase):
|
||||
self.assertEqual(patch_flags.call_count, 0)
|
||||
|
||||
@mock.patch("posthog.client.flags")
|
||||
def test_flag_with_multiple_variant_overrides(self, patch_flags):
|
||||
patch_flags.return_value = {"featureFlags": {"beta-feature": "variant-1"}}
|
||||
def test_conditions_evaluated_in_order(self, patch_flags):
|
||||
patch_flags.return_value = {"featureFlags": {"order-test": "server-variant"}}
|
||||
client = Client(FAKE_TEST_API_KEY, personal_api_key="test")
|
||||
client.feature_flags = [
|
||||
{
|
||||
"id": 1,
|
||||
"name": "Beta Feature",
|
||||
"key": "beta-feature",
|
||||
"name": "Order Test Flag",
|
||||
"key": "order-test",
|
||||
"active": True,
|
||||
"rollout_percentage": 100,
|
||||
"filters": {
|
||||
"groups": [
|
||||
{
|
||||
"rollout_percentage": 100,
|
||||
# The override applies even if the first condition matches all and gives everyone their default group
|
||||
},
|
||||
{
|
||||
"properties": [
|
||||
{
|
||||
"key": "email",
|
||||
"type": "person",
|
||||
"value": "test@posthog.com",
|
||||
"operator": "exact",
|
||||
"value": "@vip.com",
|
||||
"operator": "icontains",
|
||||
}
|
||||
],
|
||||
"rollout_percentage": 100,
|
||||
"variant": "second-variant",
|
||||
"variant": "vip-variant",
|
||||
},
|
||||
{"rollout_percentage": 50, "variant": "third-variant"},
|
||||
],
|
||||
"multivariate": {
|
||||
"variants": [
|
||||
{
|
||||
"key": "first-variant",
|
||||
"name": "First Variant",
|
||||
"rollout_percentage": 50,
|
||||
"key": "control",
|
||||
"name": "Control",
|
||||
"rollout_percentage": 100,
|
||||
},
|
||||
{
|
||||
"key": "second-variant",
|
||||
"name": "Second Variant",
|
||||
"rollout_percentage": 25,
|
||||
},
|
||||
{
|
||||
"key": "third-variant",
|
||||
"name": "Third Variant",
|
||||
"rollout_percentage": 25,
|
||||
"key": "vip-variant",
|
||||
"name": "VIP Variant",
|
||||
"rollout_percentage": 0,
|
||||
},
|
||||
]
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
self.assertEqual(
|
||||
client.get_feature_flag(
|
||||
"beta-feature",
|
||||
"test_id",
|
||||
person_properties={"email": "test@posthog.com"},
|
||||
),
|
||||
"second-variant",
|
||||
|
||||
# Even though user@vip.com would match the second condition with variant override,
|
||||
# they should match the first condition and get control
|
||||
result = client.get_feature_flag(
|
||||
"order-test",
|
||||
"user123",
|
||||
person_properties={"email": "user@vip.com"},
|
||||
)
|
||||
self.assertEqual(
|
||||
client.get_feature_flag("beta-feature", "example_id"), "third-variant"
|
||||
)
|
||||
self.assertEqual(
|
||||
client.get_feature_flag("beta-feature", "another_id"), "second-variant"
|
||||
)
|
||||
# decide not called because this can be evaluated locally
|
||||
self.assertEqual(result, "control")
|
||||
|
||||
# server not called because this can be evaluated locally
|
||||
self.assertEqual(patch_flags.call_count, 0)
|
||||
|
||||
@mock.patch("posthog.client.flags")
|
||||
@@ -3025,6 +3013,75 @@ class TestLocalEvaluation(unittest.TestCase):
|
||||
)
|
||||
self.assertEqual(patch_flags.call_count, 0)
|
||||
|
||||
@mock.patch("posthog.client.flags")
|
||||
@mock.patch("posthog.client.get")
|
||||
def test_fallback_to_api_when_flag_has_static_cohort_in_multi_condition(
|
||||
self, patch_get, patch_flags
|
||||
):
|
||||
"""
|
||||
When a flag has multiple conditions and one contains a static cohort,
|
||||
the SDK should fallback to API for the entire flag, not just skip that
|
||||
condition and evaluate the next one locally.
|
||||
|
||||
This prevents returning wrong variants when later conditions could match
|
||||
locally but the user is actually in the static cohort.
|
||||
"""
|
||||
client = Client(FAKE_TEST_API_KEY, personal_api_key=FAKE_TEST_API_KEY)
|
||||
|
||||
# Mock the local flags response - cohort 999 is NOT in cohorts map (static cohort)
|
||||
client.feature_flags = [
|
||||
{
|
||||
"id": 1,
|
||||
"key": "multi-condition-flag",
|
||||
"active": True,
|
||||
"filters": {
|
||||
"groups": [
|
||||
{
|
||||
"properties": [
|
||||
{"key": "id", "value": 999, "type": "cohort"}
|
||||
],
|
||||
"rollout_percentage": 100,
|
||||
"variant": "set-1",
|
||||
},
|
||||
{
|
||||
"properties": [
|
||||
{
|
||||
"key": "$geoip_country_code",
|
||||
"operator": "exact",
|
||||
"value": ["DE"],
|
||||
"type": "person",
|
||||
}
|
||||
],
|
||||
"rollout_percentage": 100,
|
||||
"variant": "set-8",
|
||||
},
|
||||
],
|
||||
"multivariate": {
|
||||
"variants": [
|
||||
{"key": "set-1", "rollout_percentage": 50},
|
||||
{"key": "set-8", "rollout_percentage": 50},
|
||||
]
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
client.cohorts = {} # Note: cohort 999 is NOT here - it's a static cohort
|
||||
|
||||
# Mock the API response - user is in the static cohort
|
||||
patch_flags.return_value = {"featureFlags": {"multi-condition-flag": "set-1"}}
|
||||
|
||||
result = client.get_feature_flag(
|
||||
"multi-condition-flag",
|
||||
"test-distinct-id",
|
||||
person_properties={"$geoip_country_code": "DE"},
|
||||
)
|
||||
|
||||
# Should return the API result (set-1), not local evaluation (set-8)
|
||||
self.assertEqual(result, "set-1")
|
||||
|
||||
# Verify API was called (fallback occurred)
|
||||
self.assertEqual(patch_flags.call_count, 1)
|
||||
|
||||
|
||||
class TestMatchProperties(unittest.TestCase):
|
||||
def property(self, key, value, operator=None):
|
||||
@@ -4018,6 +4075,60 @@ class TestCaptureCalls(unittest.TestCase):
|
||||
|
||||
patch_capture.reset_mock()
|
||||
|
||||
@mock.patch("posthog.client.flags")
|
||||
def test_fallback_to_api_in_get_feature_flag_payload_when_flag_has_static_cohort(
|
||||
self, patch_flags
|
||||
):
|
||||
"""
|
||||
Test that get_feature_flag_payload falls back to API when evaluating
|
||||
a flag with static cohorts, similar to get_feature_flag behavior.
|
||||
"""
|
||||
client = Client(FAKE_TEST_API_KEY, personal_api_key=FAKE_TEST_API_KEY)
|
||||
|
||||
# Mock the local flags response - cohort 999 is NOT in cohorts map (static cohort)
|
||||
client.feature_flags = [
|
||||
{
|
||||
"id": 1,
|
||||
"name": "Multi-condition Flag",
|
||||
"key": "multi-condition-flag",
|
||||
"active": True,
|
||||
"filters": {
|
||||
"groups": [
|
||||
{
|
||||
"properties": [
|
||||
{"key": "id", "value": 999, "type": "cohort"}
|
||||
],
|
||||
"rollout_percentage": 100,
|
||||
"variant": "variant-1",
|
||||
}
|
||||
],
|
||||
"multivariate": {
|
||||
"variants": [{"key": "variant-1", "rollout_percentage": 100}]
|
||||
},
|
||||
"payloads": {"variant-1": '{"message": "local-payload"}'},
|
||||
},
|
||||
}
|
||||
]
|
||||
client.cohorts = {} # Note: cohort 999 is NOT here - it's a static cohort
|
||||
|
||||
# Mock the API response - user is in the static cohort
|
||||
patch_flags.return_value = {
|
||||
"featureFlags": {"multi-condition-flag": "variant-1"},
|
||||
"featureFlagPayloads": {"multi-condition-flag": '{"message": "from-api"}'},
|
||||
}
|
||||
|
||||
# Call get_feature_flag_payload without match_value to trigger evaluation
|
||||
result = client.get_feature_flag_payload(
|
||||
"multi-condition-flag",
|
||||
"test-distinct-id",
|
||||
)
|
||||
|
||||
# Should return the API payload, not local payload
|
||||
self.assertEqual(result, {"message": "from-api"})
|
||||
|
||||
# Verify API was called (fallback occurred)
|
||||
self.assertEqual(patch_flags.call_count, 1)
|
||||
|
||||
@mock.patch.object(Client, "capture")
|
||||
@mock.patch("posthog.client.flags")
|
||||
def test_disable_geoip_get_flag_capture_call(self, patch_flags, patch_capture):
|
||||
|
||||
@@ -18,14 +18,6 @@ class TestModule(unittest.TestCase):
|
||||
"testsecret", host="http://localhost:8000", on_error=self.failed
|
||||
)
|
||||
|
||||
def test_no_api_key(self):
|
||||
self.posthog.api_key = None
|
||||
self.assertRaises(Exception, self.posthog.capture)
|
||||
|
||||
def test_no_host(self):
|
||||
self.posthog.host = None
|
||||
self.assertRaises(Exception, self.posthog.capture)
|
||||
|
||||
def test_track(self):
|
||||
res = self.posthog.capture("python module event", distinct_id="distinct_id")
|
||||
self._assert_enqueue_result(res)
|
||||
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
VERSION = "6.7.2"
|
||||
VERSION = "6.7.11"
|
||||
|
||||
if __name__ == "__main__":
|
||||
print(VERSION, end="") # noqa: T201
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Executable
+326
@@ -0,0 +1,326 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Test script to send capture_ai events to localhost:8010.
|
||||
This script tests the actual network request to a local PostHog instance.
|
||||
"""
|
||||
|
||||
from posthog import Posthog
|
||||
from uuid import uuid4
|
||||
|
||||
# Create a client pointing to localhost:8010
|
||||
posthog = Posthog(
|
||||
"test-api-key", # Use your actual project API key if needed
|
||||
host="http://localhost:8010",
|
||||
debug=True, # Enable debug mode to see detailed logs
|
||||
)
|
||||
|
||||
print("Testing capture_ai with localhost:8010")
|
||||
print("=" * 60)
|
||||
|
||||
# Test 1: $ai_generation event with blobs
|
||||
print("\n1. Testing $ai_generation event with blobs...")
|
||||
print("-" * 60)
|
||||
|
||||
trace_id = f"trace_{uuid4().hex[:8]}"
|
||||
|
||||
try:
|
||||
event_uuid = posthog.capture_ai(
|
||||
"$ai_generation",
|
||||
distinct_id="test_user_123",
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_provider": "openai",
|
||||
"$ai_trace_id": trace_id,
|
||||
"$ai_input": {
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are a helpful assistant that answers questions about Python.",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What is the difference between a list and a tuple?",
|
||||
},
|
||||
],
|
||||
"temperature": 0.7,
|
||||
"max_tokens": 500,
|
||||
},
|
||||
"$ai_output_choices": {
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "A list is mutable (can be changed) while a tuple is immutable (cannot be changed after creation). Lists use square brackets [] and tuples use parentheses ().",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"prompt_tokens": 45,
|
||||
"completion_tokens": 32,
|
||||
"total_tokens": 77,
|
||||
},
|
||||
},
|
||||
"$ai_completion_tokens": 32,
|
||||
"$ai_prompt_tokens": 45,
|
||||
"$ai_total_tokens": 77,
|
||||
"$ai_latency": 1.234,
|
||||
},
|
||||
blob_properties=["$ai_input", "$ai_output_choices"],
|
||||
)
|
||||
|
||||
if event_uuid:
|
||||
print("✓ SUCCESS: $ai_generation event sent")
|
||||
print(f" UUID: {event_uuid}")
|
||||
print(f" Trace ID: {trace_id}")
|
||||
else:
|
||||
print("✗ FAILED: No UUID returned")
|
||||
|
||||
except Exception as e:
|
||||
print(f"✗ ERROR: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
# Test 2: $ai_trace event
|
||||
print("\n2. Testing $ai_trace event...")
|
||||
print("-" * 60)
|
||||
|
||||
try:
|
||||
event_uuid = posthog.capture_ai(
|
||||
"$ai_trace",
|
||||
distinct_id="test_user_123",
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_trace_id": trace_id,
|
||||
"$ai_trace_name": "python_qa_session",
|
||||
"$ai_input_state": {
|
||||
"session_id": "session_123",
|
||||
"user_context": "learning Python",
|
||||
},
|
||||
"$ai_output_state": {"questions_answered": 1, "satisfaction_score": 5},
|
||||
},
|
||||
blob_properties=["$ai_input_state", "$ai_output_state"],
|
||||
)
|
||||
|
||||
if event_uuid:
|
||||
print("✓ SUCCESS: $ai_trace event sent")
|
||||
print(f" UUID: {event_uuid}")
|
||||
print(f" Trace ID: {trace_id}")
|
||||
else:
|
||||
print("✗ FAILED: No UUID returned")
|
||||
|
||||
except Exception as e:
|
||||
print(f"✗ ERROR: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
# Test 3: $ai_span event
|
||||
print("\n3. Testing $ai_span event...")
|
||||
print("-" * 60)
|
||||
|
||||
span_id = f"span_{uuid4().hex[:8]}"
|
||||
|
||||
try:
|
||||
event_uuid = posthog.capture_ai(
|
||||
"$ai_span",
|
||||
distinct_id="test_user_123",
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_trace_id": trace_id,
|
||||
"$ai_span_id": span_id,
|
||||
"$ai_span_name": "answer_generation",
|
||||
"$ai_parent_id": trace_id,
|
||||
"$ai_span_kind": "llm",
|
||||
"$ai_latency": 0.8,
|
||||
},
|
||||
)
|
||||
|
||||
if event_uuid:
|
||||
print("✓ SUCCESS: $ai_span event sent")
|
||||
print(f" UUID: {event_uuid}")
|
||||
print(f" Span ID: {span_id}")
|
||||
else:
|
||||
print("✗ FAILED: No UUID returned")
|
||||
|
||||
except Exception as e:
|
||||
print(f"✗ ERROR: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
# Test 4: $ai_embedding event
|
||||
print("\n4. Testing $ai_embedding event...")
|
||||
print("-" * 60)
|
||||
|
||||
try:
|
||||
event_uuid = posthog.capture_ai(
|
||||
"$ai_embedding",
|
||||
distinct_id="test_user_123",
|
||||
properties={
|
||||
"$ai_model": "text-embedding-ada-002",
|
||||
"$ai_provider": "openai",
|
||||
"$ai_trace_id": trace_id,
|
||||
"$ai_input": {
|
||||
"text": "What is the difference between a list and a tuple in Python?"
|
||||
},
|
||||
"$ai_embedding_dimension": 1536,
|
||||
"$ai_latency": 0.123,
|
||||
},
|
||||
blob_properties=["$ai_input"],
|
||||
)
|
||||
|
||||
if event_uuid:
|
||||
print("✓ SUCCESS: $ai_embedding event sent")
|
||||
print(f" UUID: {event_uuid}")
|
||||
else:
|
||||
print("✗ FAILED: No UUID returned")
|
||||
|
||||
except Exception as e:
|
||||
print(f"✗ ERROR: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
# Test 5: $ai_metric event
|
||||
print("\n5. Testing $ai_metric event...")
|
||||
print("-" * 60)
|
||||
|
||||
try:
|
||||
event_uuid = posthog.capture_ai(
|
||||
"$ai_metric",
|
||||
distinct_id="test_user_123",
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_trace_id": trace_id,
|
||||
"$ai_metric_name": "response_quality",
|
||||
"$ai_metric_value": "0.95",
|
||||
},
|
||||
)
|
||||
|
||||
if event_uuid:
|
||||
print("✓ SUCCESS: $ai_metric event sent")
|
||||
print(f" UUID: {event_uuid}")
|
||||
else:
|
||||
print("✗ FAILED: No UUID returned")
|
||||
|
||||
except Exception as e:
|
||||
print(f"✗ ERROR: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
# Test 6: $ai_feedback event
|
||||
print("\n6. Testing $ai_feedback event...")
|
||||
print("-" * 60)
|
||||
|
||||
try:
|
||||
event_uuid = posthog.capture_ai(
|
||||
"$ai_feedback",
|
||||
distinct_id="test_user_123",
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_trace_id": trace_id,
|
||||
"$ai_feedback_text": "Great explanation! Very clear and helpful.",
|
||||
"$ai_feedback_rating": 5,
|
||||
},
|
||||
)
|
||||
|
||||
if event_uuid:
|
||||
print("✓ SUCCESS: $ai_feedback event sent")
|
||||
print(f" UUID: {event_uuid}")
|
||||
else:
|
||||
print("✗ FAILED: No UUID returned")
|
||||
|
||||
except Exception as e:
|
||||
print(f"✗ ERROR: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
# Test 7: Test with custom blob properties
|
||||
print("\n7. Testing with custom blob properties...")
|
||||
print("-" * 60)
|
||||
|
||||
try:
|
||||
event_uuid = posthog.capture_ai(
|
||||
"$ai_generation",
|
||||
distinct_id="test_user_123",
|
||||
properties={
|
||||
"$ai_model": "claude-3-opus",
|
||||
"$ai_provider": "anthropic",
|
||||
"$ai_trace_id": trace_id,
|
||||
"$ai_input": {
|
||||
"messages": [{"role": "user", "content": "Write a haiku about Python"}]
|
||||
},
|
||||
"$ai_output_choices": {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Snake glides through code\nSimple syntax, powerful tools\nDevelopers smile",
|
||||
}
|
||||
}
|
||||
]
|
||||
},
|
||||
"$ai_custom_data": {
|
||||
"large_context": "This is some large custom data that should be sent as a blob"
|
||||
},
|
||||
"$ai_completion_tokens": 20,
|
||||
"$ai_prompt_tokens": 10,
|
||||
},
|
||||
# Custom blob properties - including the default ones plus a custom one
|
||||
blob_properties=["$ai_input", "$ai_output_choices", "$ai_custom_data"],
|
||||
)
|
||||
|
||||
if event_uuid:
|
||||
print("✓ SUCCESS: Event with custom blob properties sent")
|
||||
print(f" UUID: {event_uuid}")
|
||||
else:
|
||||
print("✗ FAILED: No UUID returned")
|
||||
|
||||
except Exception as e:
|
||||
print(f"✗ ERROR: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
# Test 8: Test with groups
|
||||
print("\n8. Testing with groups...")
|
||||
print("-" * 60)
|
||||
|
||||
try:
|
||||
event_uuid = posthog.capture_ai(
|
||||
"$ai_generation",
|
||||
distinct_id="test_user_123",
|
||||
groups={"company": "posthog_inc", "team": "engineering"},
|
||||
properties={
|
||||
"$ai_model": "gpt-4",
|
||||
"$ai_provider": "openai",
|
||||
"$ai_trace_id": trace_id,
|
||||
"$ai_input": {"messages": [{"role": "user", "content": "test"}]},
|
||||
"$ai_output_choices": {
|
||||
"choices": [{"message": {"role": "assistant", "content": "response"}}]
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
if event_uuid:
|
||||
print("✓ SUCCESS: Event with groups sent")
|
||||
print(f" UUID: {event_uuid}")
|
||||
else:
|
||||
print("✗ FAILED: No UUID returned")
|
||||
|
||||
except Exception as e:
|
||||
print(f"✗ ERROR: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print("All tests completed!")
|
||||
print("\nMake sure your local PostHog instance is running on http://localhost:8010")
|
||||
print("and that the /i/v0/ai endpoint is available.")
|
||||
Reference in New Issue
Block a user