Compare commits

...
5 Commits
Author SHA1 Message Date
Radu RaiceaandGitHub f5719f39da fix(llma): default prompts url (#423) 2026-02-04 15:10:00 +00:00
Radu RaiceaandGitHub d4f2d6dfb0 fix(llma): small fixes for prompt management (#420)
* fix(llma): small fixes for prompt management

* fix(llma): tests

* fix(llma): tests
2026-02-04 09:49:19 +02:00
José SequeiraandGitHub 72f448816c feat: SDK Compliance (#397)
* feat: SDK Compliance
2026-01-30 16:12:43 +01:00
Radu RaiceaandGitHub 4350389f93 feat(llma): add prompt management (#417)
* feat(llma): add prompt management

* chore(llma): bump version

* fix(llma): use SDK session with retry logic for prompt fetching

Use _get_session() from posthog/request.py instead of raw requests.get()
to benefit from the SDK's existing retry configuration on transient
network failures.
2026-01-30 08:43:04 -05:00
c32c78312f feat(llma): pass raw provider usage metadata for backend cost calculations (#411)
* feat: pass raw provider usage metadata for backend cost calculations

Add raw_usage field to TokenUsage type to capture raw provider usage metadata (OpenAI, Anthropic, Gemini). This enables the backend to extract modality-specific token counts (text vs image vs audio) for accurate cost calculations.

- Add raw_usage field to TokenUsage TypedDict
- Update all provider converters to capture raw usage:
  - OpenAI: capture response.usage and chunk usage
  - Anthropic: capture usage from message_start and message_delta events
  - Gemini: capture usage_metadata from responses and chunks
- Pass raw usage as $ai_usage property in PostHog events
- Update merge_usage_stats to handle raw_usage in both modes
- Add tests verifying $ai_usage is captured for all providers

Backend will extract provider-specific details and delete $ai_usage after processing to avoid bloating properties.

Co-Authored-By: Claude Sonnet 4.5 <noreply@anthropic.com>

* fix: add serialize_raw_usage helper to ensure JSON serializability

Address PR review feedback from @andrewm4894:

1. **Serialization**: Add serialize_raw_usage() helper with fallback chain:
   - .model_dump() for Pydantic models (OpenAI/Anthropic)
   - .to_dict() for protobuf-like objects
   - vars() for simple objects
   - str() as last resort
   This ensures we never pass unserializable objects to PostHog client.

2. **Data loss prevention**: Change from replacing to merging raw_usage in
   incremental mode. For Anthropic streaming, message_start has input token
   details and message_delta has output token details - merging preserves
   both instead of losing input data.

3. **Test coverage**: Enhanced tests to verify:
   - JSON serializability with json.dumps()
   - Expected structure of raw_usage dicts
   - Coverage for both non-streaming and streaming modes
   - Fixed Gemini test mocks to return proper dicts from model_dump()

Co-Authored-By: Claude Sonnet 4.5 <noreply@anthropic.com>

* refactor: move raw_usage serialization from utils to converters

Address PR feedback from @andrewm4894 - serialize in converters, not utils.

**Problem:**
Utils was receiving raw Pydantic/protobuf objects and serializing them,
which meant provider-specific knowledge leaked into generic code.

**Solution:**
Move serialization into converters where provider context exists:

Converters (NEW):
- OpenAI: serialize_raw_usage(response.usage) → dict
- Anthropic: serialize_raw_usage(event.usage) → dict
- Gemini: serialize_raw_usage(metadata) → dict

Utils (SIMPLIFIED):
- Just passes dicts through, no serialization needed
- Merge operations work with dicts only

**Benefits:**
1. Type correctness: raw_usage is always Dict[str, Any]
2. Separation of concerns: converters handle provider formats
3. Fail fast: serialization errors in converters with context
4. Cleaner abstraction: utils doesn't know about Pydantic/protobuf

**Flow:**
Provider object → Converter serializes → dict → Utils → PostHog

Co-Authored-By: Claude Sonnet 4.5 <noreply@anthropic.com>

* fix: add type annotation for current_raw to satisfy mypy

Fix mypy error: "Need type annotation for 'current_raw'"

Extract value first, then apply explicit type annotation with ternary
conditional to satisfy mypy's type checker.

Co-Authored-By: Claude Sonnet 4.5 <noreply@anthropic.com>

---------

Co-authored-by: Claude Sonnet 4.5 <noreply@anthropic.com>
2026-01-28 11:51:02 +02:00
19 changed files with 1615 additions and 1 deletions
+21
View File
@@ -0,0 +1,21 @@
name: SDK Compliance Tests
permissions:
contents: read
packages: read
pull-requests: write
on:
pull_request:
push:
branches:
- master
jobs:
compliance:
name: PostHog SDK compliance tests
uses: PostHog/posthog-sdk-test-harness/.github/workflows/test-sdk-action.yml@main
with:
adapter-dockerfile: "sdk_compliance_adapter/Dockerfile"
adapter-context: "."
test-harness-version: "latest"
+14
View File
@@ -1,3 +1,17 @@
# 7.8.2 - 2026-02-04
fix(llma): fix prompts default url
# 7.8.1 - 2026-02-03
fix(llma): small fixes for prompt management
# 7.8.0 - 2026-01-28
feat(llma): add prompt management
Adds the Prompt Management feature. At the time of release, this feature is in a closed alpha.
# 7.7.0 - 2026-01-15
feat(ai): Add OpenAI Agents SDK integration
+3
View File
@@ -0,0 +1,3 @@
from posthog.ai.prompts import Prompts
__all__ = ["Prompts"]
@@ -17,6 +17,7 @@ from posthog.ai.types import (
TokenUsage,
ToolInProgress,
)
from posthog.ai.utils import serialize_raw_usage
def format_anthropic_response(response: Any) -> List[FormattedMessage]:
@@ -221,6 +222,12 @@ def extract_anthropic_usage_from_response(response: Any) -> TokenUsage:
if web_search_count > 0:
result["web_search_count"] = web_search_count
# Capture raw usage metadata for backend processing
# Serialize to dict here in the converter (not in utils)
serialized = serialize_raw_usage(response.usage)
if serialized:
result["raw_usage"] = serialized
return result
@@ -247,6 +254,11 @@ def extract_anthropic_usage_from_event(event: Any) -> TokenUsage:
usage["cache_read_input_tokens"] = getattr(
event.message.usage, "cache_read_input_tokens", 0
)
# Capture raw usage metadata for backend processing
# Serialize to dict here in the converter (not in utils)
serialized = serialize_raw_usage(event.message.usage)
if serialized:
usage["raw_usage"] = serialized
# Handle usage stats from message_delta event
if hasattr(event, "usage") and event.usage:
@@ -262,6 +274,12 @@ def extract_anthropic_usage_from_event(event: Any) -> TokenUsage:
if web_search_count > 0:
usage["web_search_count"] = web_search_count
# Capture raw usage metadata for backend processing
# Serialize to dict here in the converter (not in utils)
serialized = serialize_raw_usage(event.usage)
if serialized:
usage["raw_usage"] = serialized
return usage
+7
View File
@@ -12,6 +12,7 @@ from posthog.ai.types import (
FormattedMessage,
TokenUsage,
)
from posthog.ai.utils import serialize_raw_usage
class GeminiPart(TypedDict, total=False):
@@ -487,6 +488,12 @@ def _extract_usage_from_metadata(metadata: Any) -> TokenUsage:
if reasoning_tokens and reasoning_tokens > 0:
usage["reasoning_tokens"] = reasoning_tokens
# Capture raw usage metadata for backend processing
# Serialize to dict here in the converter (not in utils)
serialized = serialize_raw_usage(metadata)
if serialized:
usage["raw_usage"] = serialized
return usage
+19
View File
@@ -16,6 +16,7 @@ from posthog.ai.types import (
FormattedTextContent,
TokenUsage,
)
from posthog.ai.utils import serialize_raw_usage
def format_openai_response(response: Any) -> List[FormattedMessage]:
@@ -429,6 +430,12 @@ def extract_openai_usage_from_response(response: Any) -> TokenUsage:
if web_search_count > 0:
result["web_search_count"] = web_search_count
# Capture raw usage metadata for backend processing
# Serialize to dict here in the converter (not in utils)
serialized = serialize_raw_usage(response.usage)
if serialized:
result["raw_usage"] = serialized
return result
@@ -482,6 +489,12 @@ def extract_openai_usage_from_chunk(
chunk.usage.completion_tokens_details.reasoning_tokens
)
# Capture raw usage metadata for backend processing
# Serialize to dict here in the converter (not in utils)
serialized = serialize_raw_usage(chunk.usage)
if serialized:
usage["raw_usage"] = serialized
elif provider_type == "responses":
# For Responses API, usage is only in chunk.response.usage for completed events
if hasattr(chunk, "type") and chunk.type == "response.completed":
@@ -516,6 +529,12 @@ def extract_openai_usage_from_chunk(
if web_search_count > 0:
usage["web_search_count"] = web_search_count
# Capture raw usage metadata for backend processing
# Serialize to dict here in the converter (not in utils)
serialized = serialize_raw_usage(response_usage)
if serialized:
usage["raw_usage"] = serialized
return usage
+272
View File
@@ -0,0 +1,272 @@
"""
Prompt management for PostHog AI SDK.
Fetch and compile LLM prompts from PostHog with caching and fallback support.
"""
import logging
import re
import time
import urllib.parse
from typing import Any, Dict, Optional, Union
from posthog.request import USER_AGENT, _get_session
from posthog.utils import remove_trailing_slash
log = logging.getLogger("posthog")
APP_ENDPOINT = "https://us.posthog.com"
DEFAULT_CACHE_TTL_SECONDS = 300 # 5 minutes
PromptVariables = Dict[str, Union[str, int, float, bool]]
class CachedPrompt:
"""Cached prompt with metadata."""
def __init__(self, prompt: str, fetched_at: float):
self.prompt = prompt
self.fetched_at = fetched_at
def _is_prompt_api_response(data: Any) -> bool:
"""Check if the response is a valid prompt API response."""
return (
isinstance(data, dict)
and "prompt" in data
and isinstance(data.get("prompt"), str)
)
class Prompts:
"""
Fetch and compile LLM prompts from PostHog.
Can be initialized with a PostHog client or with direct options.
Examples:
```python
from posthog import Posthog
from posthog.ai.prompts import Prompts
# With PostHog client
posthog = Posthog('phc_xxx', host='https://us.posthog.com', personal_api_key='phx_xxx')
prompts = Prompts(posthog)
# Or with direct options (no PostHog client needed)
prompts = Prompts(personal_api_key='phx_xxx', host='https://us.posthog.com')
# Fetch with caching and fallback
template = prompts.get('support-system-prompt', fallback='You are a helpful assistant.')
# Compile with variables
system_prompt = prompts.compile(template, {
'company': 'Acme Corp',
'tier': 'premium',
})
```
"""
def __init__(
self,
posthog: Optional[Any] = None,
*,
personal_api_key: Optional[str] = None,
host: Optional[str] = None,
default_cache_ttl_seconds: Optional[int] = None,
):
"""
Initialize Prompts.
Args:
posthog: PostHog client instance (optional if personal_api_key provided)
personal_api_key: Direct API key (optional if posthog provided)
host: PostHog host (defaults to app endpoint)
default_cache_ttl_seconds: Default cache TTL (defaults to 300)
"""
self._default_cache_ttl_seconds = (
default_cache_ttl_seconds or DEFAULT_CACHE_TTL_SECONDS
)
self._cache: Dict[str, CachedPrompt] = {}
if posthog is not None:
self._personal_api_key = getattr(posthog, "personal_api_key", None) or ""
self._host = remove_trailing_slash(
getattr(posthog, "raw_host", None) or APP_ENDPOINT
)
else:
self._personal_api_key = personal_api_key or ""
self._host = remove_trailing_slash(host or APP_ENDPOINT)
def get(
self,
name: str,
*,
cache_ttl_seconds: Optional[int] = None,
fallback: Optional[str] = None,
) -> str:
"""
Fetch a prompt by name from the PostHog API.
Caching behavior:
1. If cache is fresh, return cached value
2. If fetch fails and cache exists (stale), return stale cache with warning
3. If fetch fails and fallback provided, return fallback with warning
4. If fetch fails with no cache/fallback, raise exception
Args:
name: The name of the prompt to fetch
cache_ttl_seconds: Cache TTL in seconds (defaults to instance default)
fallback: Fallback prompt to use if fetch fails and no cache available
Returns:
The prompt string
Raises:
Exception: If the prompt cannot be fetched and no fallback is available
"""
ttl = (
cache_ttl_seconds
if cache_ttl_seconds is not None
else self._default_cache_ttl_seconds
)
# Check cache first
cached = self._cache.get(name)
now = time.time()
if cached is not None:
is_fresh = (now - cached.fetched_at) < ttl
if is_fresh:
return cached.prompt
# Try to fetch from API
try:
prompt = self._fetch_prompt_from_api(name)
fetched_at = time.time()
# Update cache
self._cache[name] = CachedPrompt(prompt=prompt, fetched_at=fetched_at)
return prompt
except Exception as error:
# Fallback order:
# 1. Return stale cache (with warning)
if cached is not None:
log.warning(
'[PostHog Prompts] Failed to fetch prompt "%s", using stale cache: %s',
name,
error,
)
return cached.prompt
# 2. Return fallback (with warning)
if fallback is not None:
log.warning(
'[PostHog Prompts] Failed to fetch prompt "%s", using fallback: %s',
name,
error,
)
return fallback
# 3. Raise error
raise
def compile(self, prompt: str, variables: PromptVariables) -> str:
"""
Replace {{variableName}} placeholders with values.
Unmatched variables are left unchanged.
Supports variable names with hyphens and dots (e.g., user-id, company.name).
Args:
prompt: The prompt template string
variables: Object containing variable values
Returns:
The compiled prompt string
"""
def replace_variable(match: re.Match) -> str:
variable_name = match.group(1)
if variable_name in variables:
return str(variables[variable_name])
return match.group(0)
return re.sub(r"\{\{([\w.-]+)\}\}", replace_variable, prompt)
def clear_cache(self, name: Optional[str] = None) -> None:
"""
Clear cached prompts.
Args:
name: Specific prompt to clear. If None, clears all cached prompts.
"""
if name is not None:
self._cache.pop(name, None)
else:
self._cache.clear()
def _fetch_prompt_from_api(self, name: str) -> str:
"""
Fetch prompt from PostHog API.
Endpoint: {host}/api/environments/@current/llm_prompts/name/{encoded_name}/
Auth: Bearer {personal_api_key}
Args:
name: The name of the prompt to fetch
Returns:
The prompt string
Raises:
Exception: If the prompt cannot be fetched
"""
if not self._personal_api_key:
raise Exception(
"[PostHog Prompts] personal_api_key is required to fetch prompts. "
"Please provide it when initializing the Prompts instance."
)
encoded_name = urllib.parse.quote(name, safe="")
url = f"{self._host}/api/environments/@current/llm_prompts/name/{encoded_name}/"
headers = {
"Authorization": f"Bearer {self._personal_api_key}",
"User-Agent": USER_AGENT,
}
response = _get_session().get(url, headers=headers, timeout=10)
if not response.ok:
if response.status_code == 404:
raise Exception(f'[PostHog Prompts] Prompt "{name}" not found')
if response.status_code == 403:
raise Exception(
f'[PostHog Prompts] Access denied for prompt "{name}". '
"Check that your personal_api_key has the correct permissions and the LLM prompts feature is enabled."
)
raise Exception(
f'[PostHog Prompts] Failed to fetch prompt "{name}": HTTP {response.status_code}'
)
try:
data = response.json()
except Exception:
raise Exception(
f'[PostHog Prompts] Invalid response format for prompt "{name}"'
)
if not _is_prompt_api_response(data):
raise Exception(
f'[PostHog Prompts] Invalid response format for prompt "{name}"'
)
return data["prompt"]
+1
View File
@@ -64,6 +64,7 @@ class TokenUsage(TypedDict, total=False):
cache_creation_input_tokens: Optional[int]
reasoning_tokens: Optional[int]
web_search_count: Optional[int]
raw_usage: Optional[Any] # Raw provider usage metadata for backend processing
class ProviderResponse(TypedDict, total=False):
+78
View File
@@ -13,6 +13,54 @@ from posthog.ai.types import FormattedMessage, StreamingEventData, TokenUsage
from posthog.client import Client as PostHogClient
def serialize_raw_usage(raw_usage: Any) -> Optional[Dict[str, Any]]:
"""
Convert raw provider usage objects to JSON-serializable dicts.
Handles Pydantic models (OpenAI/Anthropic) and protobuf-like objects (Gemini)
with a fallback chain to ensure we never pass unserializable objects to PostHog.
Args:
raw_usage: Raw usage object from provider SDK
Returns:
Plain dict or None if conversion fails
"""
if raw_usage is None:
return None
# Already a dict
if isinstance(raw_usage, dict):
return raw_usage
# Try Pydantic model_dump() (OpenAI/Anthropic)
if hasattr(raw_usage, "model_dump") and callable(raw_usage.model_dump):
try:
return raw_usage.model_dump()
except Exception:
pass
# Try to_dict() (some protobuf objects)
if hasattr(raw_usage, "to_dict") and callable(raw_usage.to_dict):
try:
return raw_usage.to_dict()
except Exception:
pass
# Try __dict__ / vars() for simple objects
try:
return vars(raw_usage)
except Exception:
pass
# Last resort: convert to string representation
# This ensures we always return something rather than failing
try:
return {"_raw": str(raw_usage)}
except Exception:
return None
def merge_usage_stats(
target: TokenUsage, source: TokenUsage, mode: str = "incremental"
) -> None:
@@ -60,6 +108,17 @@ def merge_usage_stats(
current = target.get("web_search_count") or 0
target["web_search_count"] = max(current, source_web_search)
# Merge raw_usage to avoid losing data from earlier events
# For Anthropic streaming: message_start has input tokens, message_delta has output
# Note: raw_usage is already serialized by converters, so it's a dict
source_raw_usage = source.get("raw_usage")
if source_raw_usage is not None and isinstance(source_raw_usage, dict):
current_raw_value = target.get("raw_usage")
current_raw: Dict[str, Any] = (
current_raw_value if isinstance(current_raw_value, dict) else {}
)
target["raw_usage"] = {**current_raw, **source_raw_usage}
elif mode == "cumulative":
# Replace with latest values (already cumulative)
if source.get("input_tokens") is not None:
@@ -76,6 +135,9 @@ def merge_usage_stats(
target["reasoning_tokens"] = source["reasoning_tokens"]
if source.get("web_search_count") is not None:
target["web_search_count"] = source["web_search_count"]
# Note: raw_usage is already serialized by converters, so it's a dict
if source.get("raw_usage") is not None:
target["raw_usage"] = source["raw_usage"]
else:
raise ValueError(f"Invalid mode: {mode}. Must be 'incremental' or 'cumulative'")
@@ -332,6 +394,11 @@ def call_llm_and_track_usage(
if web_search_count is not None and web_search_count > 0:
tag("$ai_web_search_count", web_search_count)
raw_usage = usage.get("raw_usage")
if raw_usage is not None:
# Already serialized by converters
tag("$ai_usage", raw_usage)
if posthog_distinct_id is None:
tag("$process_person_profile", False)
@@ -457,6 +524,11 @@ async def call_llm_and_track_usage_async(
if web_search_count is not None and web_search_count > 0:
tag("$ai_web_search_count", web_search_count)
raw_usage = usage.get("raw_usage")
if raw_usage is not None:
# Already serialized by converters
tag("$ai_usage", raw_usage)
if posthog_distinct_id is None:
tag("$process_person_profile", False)
@@ -594,6 +666,12 @@ def capture_streaming_event(
):
event_properties["$ai_web_search_count"] = web_search_count
# Add raw usage metadata if present (all providers)
raw_usage = event_data["usage_stats"].get("raw_usage")
if raw_usage is not None:
# Already serialized by converters
event_properties["$ai_usage"] = raw_usage
# Handle provider-specific fields
if (
event_data["provider"] == "openai"
@@ -1,3 +1,4 @@
import json
from unittest.mock import patch
import pytest
@@ -306,6 +307,15 @@ def test_basic_completion(mock_client, mock_anthropic_response):
assert props["$ai_http_status"] == 200
assert props["foo"] == "bar"
assert isinstance(props["$ai_latency"], float)
# Verify raw usage metadata is passed for backend processing
assert "$ai_usage" in props
assert props["$ai_usage"] is not None
# Verify it's JSON-serializable
json.dumps(props["$ai_usage"])
# Verify it has expected structure
assert isinstance(props["$ai_usage"], dict)
assert "input_tokens" in props["$ai_usage"]
assert "output_tokens" in props["$ai_usage"]
def test_groups(mock_client, mock_anthropic_response):
@@ -918,6 +928,16 @@ def test_streaming_with_tool_calls(mock_client, mock_anthropic_stream_with_tools
assert props["$ai_cache_read_input_tokens"] == 5
assert props["$ai_cache_creation_input_tokens"] == 0
# Verify raw usage is captured in streaming mode (merged from events)
assert "$ai_usage" in props
assert props["$ai_usage"] is not None
# Verify it's JSON-serializable
json.dumps(props["$ai_usage"])
# Verify it has expected structure (merged from message_start and message_delta)
assert isinstance(props["$ai_usage"], dict)
assert "input_tokens" in props["$ai_usage"]
assert "output_tokens" in props["$ai_usage"]
def test_async_streaming_with_tool_calls(mock_client, mock_anthropic_stream_with_tools):
"""Test that tool calls are properly captured in async streaming mode."""
+55
View File
@@ -1,3 +1,4 @@
import json
from unittest.mock import MagicMock, patch
import pytest
@@ -34,6 +35,13 @@ def mock_gemini_response():
# Ensure cache and reasoning tokens are not present (not MagicMock)
mock_usage.cached_content_token_count = 0
mock_usage.thoughts_token_count = 0
# Make model_dump() return a proper dict for serialization
mock_usage.model_dump.return_value = {
"prompt_token_count": 20,
"candidates_token_count": 10,
"cached_content_token_count": 0,
"thoughts_token_count": 0,
}
mock_response.usage_metadata = mock_usage
mock_candidate = MagicMock()
@@ -69,6 +77,13 @@ def mock_gemini_response_with_function_calls():
mock_usage.candidates_token_count = 15
mock_usage.cached_content_token_count = 0
mock_usage.thoughts_token_count = 0
# Make model_dump() return a proper dict for serialization
mock_usage.model_dump.return_value = {
"prompt_token_count": 25,
"candidates_token_count": 15,
"cached_content_token_count": 0,
"thoughts_token_count": 0,
}
mock_response.usage_metadata = mock_usage
# Mock function call
@@ -117,6 +132,13 @@ def mock_gemini_response_function_calls_only():
mock_usage.candidates_token_count = 12
mock_usage.cached_content_token_count = 0
mock_usage.thoughts_token_count = 0
# Make model_dump() return a proper dict for serialization
mock_usage.model_dump.return_value = {
"prompt_token_count": 30,
"candidates_token_count": 12,
"cached_content_token_count": 0,
"thoughts_token_count": 0,
}
mock_response.usage_metadata = mock_usage
# Mock function call
@@ -174,6 +196,15 @@ def test_new_client_basic_generation(
assert props["foo"] == "bar"
assert "$ai_trace_id" in props
assert props["$ai_latency"] > 0
# Verify raw usage metadata is passed for backend processing
assert "$ai_usage" in props
assert props["$ai_usage"] is not None
# Verify it's JSON-serializable
json.dumps(props["$ai_usage"])
# Verify it has expected structure
assert isinstance(props["$ai_usage"], dict)
assert "prompt_token_count" in props["$ai_usage"]
assert "candidates_token_count" in props["$ai_usage"]
def test_new_client_streaming_with_generate_content_stream(
@@ -810,6 +841,13 @@ def test_streaming_cache_and_reasoning_tokens(mock_client, mock_google_genai_cli
chunk1_usage.candidates_token_count = 5
chunk1_usage.cached_content_token_count = 30 # Cache tokens
chunk1_usage.thoughts_token_count = 0
# Make model_dump() return a proper dict for serialization
chunk1_usage.model_dump.return_value = {
"prompt_token_count": 100,
"candidates_token_count": 5,
"cached_content_token_count": 30,
"thoughts_token_count": 0,
}
chunk1.usage_metadata = chunk1_usage
chunk2 = MagicMock()
@@ -819,6 +857,13 @@ def test_streaming_cache_and_reasoning_tokens(mock_client, mock_google_genai_cli
chunk2_usage.candidates_token_count = 10
chunk2_usage.cached_content_token_count = 30 # Same cache tokens
chunk2_usage.thoughts_token_count = 5 # Reasoning tokens
# Make model_dump() return a proper dict for serialization
chunk2_usage.model_dump.return_value = {
"prompt_token_count": 100,
"candidates_token_count": 10,
"cached_content_token_count": 30,
"thoughts_token_count": 5,
}
chunk2.usage_metadata = chunk2_usage
mock_stream = iter([chunk1, chunk2])
@@ -848,6 +893,16 @@ def test_streaming_cache_and_reasoning_tokens(mock_client, mock_google_genai_cli
assert props["$ai_cache_read_input_tokens"] == 30
assert props["$ai_reasoning_tokens"] == 5
# Verify raw usage is captured in streaming mode (merged from chunks)
assert "$ai_usage" in props
assert props["$ai_usage"] is not None
# Verify it's JSON-serializable
json.dumps(props["$ai_usage"])
# Verify it has expected structure
assert isinstance(props["$ai_usage"], dict)
assert "prompt_token_count" in props["$ai_usage"]
assert "candidates_token_count" in props["$ai_usage"]
def test_web_search_grounding(mock_client, mock_google_genai_client):
"""Test web search detection via grounding_metadata."""
+20
View File
@@ -1,3 +1,4 @@
import json
import time
from unittest.mock import AsyncMock, patch
@@ -496,6 +497,15 @@ def test_basic_completion(mock_client, mock_openai_response):
assert props["$ai_http_status"] == 200
assert props["foo"] == "bar"
assert isinstance(props["$ai_latency"], float)
# Verify raw usage metadata is passed for backend processing
assert "$ai_usage" in props
assert props["$ai_usage"] is not None
# Verify it's JSON-serializable
json.dumps(props["$ai_usage"])
# Verify it has expected structure
assert isinstance(props["$ai_usage"], dict)
assert "prompt_tokens" in props["$ai_usage"]
assert "completion_tokens" in props["$ai_usage"]
def test_embeddings(mock_client, mock_embedding_response):
@@ -922,6 +932,16 @@ def test_streaming_with_tool_calls(mock_client, streaming_tool_call_chunks):
assert props["$ai_input_tokens"] == 20
assert props["$ai_output_tokens"] == 15
# Verify raw usage is captured in streaming mode
assert "$ai_usage" in props
assert props["$ai_usage"] is not None
# Verify it's JSON-serializable
json.dumps(props["$ai_usage"])
# Verify it has expected structure (merged from chunks)
assert isinstance(props["$ai_usage"], dict)
assert "prompt_tokens" in props["$ai_usage"]
assert "completion_tokens" in props["$ai_usage"]
# test responses api
def test_responses_api(mock_client, mock_openai_response_with_responses_api):
+577
View File
@@ -0,0 +1,577 @@
import unittest
from unittest.mock import MagicMock, patch
from posthog.ai.prompts import Prompts
class MockResponse:
"""Mock HTTP response for testing."""
def __init__(self, json_data=None, status_code=200, ok=True):
self._json_data = json_data
self.status_code = status_code
self.ok = ok
def json(self):
if self._json_data is None:
raise ValueError("No JSON data")
return self._json_data
class TestPrompts(unittest.TestCase):
"""Tests for the Prompts class."""
mock_prompt_response = {
"id": 1,
"name": "test-prompt",
"prompt": "Hello, {{name}}! You are a helpful assistant for {{company}}.",
"version": 1,
"created_by": "user@example.com",
"created_at": "2024-01-01T00:00:00Z",
"updated_at": "2024-01-01T00:00:00Z",
"deleted": False,
}
def create_mock_posthog(
self, personal_api_key="phx_test_key", host="https://us.posthog.com"
):
"""Create a mock PostHog client."""
mock = MagicMock()
mock.personal_api_key = personal_api_key
mock.raw_host = host
return mock
class TestPromptsGet(TestPrompts):
"""Tests for the Prompts.get() method."""
@patch("posthog.ai.prompts._get_session")
def test_successfully_fetch_a_prompt(self, mock_get_session):
"""Should successfully fetch a prompt."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(json_data=self.mock_prompt_response)
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
result = prompts.get("test-prompt")
self.assertEqual(result, self.mock_prompt_response["prompt"])
mock_get.assert_called_once()
call_args = mock_get.call_args
self.assertEqual(
call_args[0][0],
"https://us.posthog.com/api/environments/@current/llm_prompts/name/test-prompt/",
)
self.assertIn("Authorization", call_args[1]["headers"])
self.assertEqual(
call_args[1]["headers"]["Authorization"], "Bearer phx_test_key"
)
@patch("posthog.ai.prompts._get_session")
@patch("posthog.ai.prompts.time.time")
def test_return_cached_prompt_when_fresh(self, mock_time, mock_get_session):
"""Should return cached prompt when fresh (no API call)."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(json_data=self.mock_prompt_response)
mock_time.return_value = 1000.0
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
# First call - fetches from API
result1 = prompts.get("test-prompt", cache_ttl_seconds=300)
self.assertEqual(result1, self.mock_prompt_response["prompt"])
self.assertEqual(mock_get.call_count, 1)
# Advance time by 60 seconds (still within TTL)
mock_time.return_value = 1060.0
# Second call - should use cache
result2 = prompts.get("test-prompt", cache_ttl_seconds=300)
self.assertEqual(result2, self.mock_prompt_response["prompt"])
self.assertEqual(mock_get.call_count, 1) # No additional fetch
@patch("posthog.ai.prompts._get_session")
@patch("posthog.ai.prompts.time.time")
def test_refetch_when_cache_is_stale(self, mock_time, mock_get_session):
"""Should refetch when cache is stale."""
mock_get = mock_get_session.return_value.get
updated_prompt_response = {
**self.mock_prompt_response,
"prompt": "Updated prompt: Hello, {{name}}!",
}
mock_get.side_effect = [
MockResponse(json_data=self.mock_prompt_response),
MockResponse(json_data=updated_prompt_response),
]
mock_time.return_value = 1000.0
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
# First call - fetches from API
result1 = prompts.get("test-prompt", cache_ttl_seconds=60)
self.assertEqual(result1, self.mock_prompt_response["prompt"])
self.assertEqual(mock_get.call_count, 1)
# Advance time past TTL
mock_time.return_value = 1061.0
# Second call - should refetch
result2 = prompts.get("test-prompt", cache_ttl_seconds=60)
self.assertEqual(result2, updated_prompt_response["prompt"])
self.assertEqual(mock_get.call_count, 2)
@patch("posthog.ai.prompts._get_session")
@patch("posthog.ai.prompts.time.time")
@patch("posthog.ai.prompts.log")
def test_use_stale_cache_on_fetch_failure_with_warning(
self, mock_log, mock_time, mock_get_session
):
"""Should use stale cache on fetch failure with warning."""
mock_get = mock_get_session.return_value.get
mock_get.side_effect = [
MockResponse(json_data=self.mock_prompt_response),
Exception("Network error"),
]
mock_time.return_value = 1000.0
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
# First call - populates cache
result1 = prompts.get("test-prompt", cache_ttl_seconds=60)
self.assertEqual(result1, self.mock_prompt_response["prompt"])
# Advance time past TTL
mock_time.return_value = 1061.0
# Second call - should use stale cache
result2 = prompts.get("test-prompt", cache_ttl_seconds=60)
self.assertEqual(result2, self.mock_prompt_response["prompt"])
# Check warning was logged
mock_log.warning.assert_called()
warning_call = mock_log.warning.call_args
self.assertIn("using stale cache", warning_call[0][0])
@patch("posthog.ai.prompts._get_session")
@patch("posthog.ai.prompts.log")
def test_use_fallback_when_no_cache_and_fetch_fails_with_warning(
self, mock_log, mock_get_session
):
"""Should use fallback when no cache and fetch fails with warning."""
mock_get = mock_get_session.return_value.get
mock_get.side_effect = Exception("Network error")
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
fallback = "Default system prompt."
result = prompts.get("test-prompt", fallback=fallback)
self.assertEqual(result, fallback)
# Check warning was logged
mock_log.warning.assert_called()
warning_call = mock_log.warning.call_args
self.assertIn("using fallback", warning_call[0][0])
@patch("posthog.ai.prompts._get_session")
def test_throw_when_no_cache_no_fallback_and_fetch_fails(self, mock_get_session):
"""Should throw when no cache, no fallback, and fetch fails."""
mock_get = mock_get_session.return_value.get
mock_get.side_effect = Exception("Network error")
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
with self.assertRaises(Exception) as context:
prompts.get("test-prompt")
self.assertIn("Network error", str(context.exception))
@patch("posthog.ai.prompts._get_session")
def test_handle_404_response(self, mock_get_session):
"""Should handle 404 response."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(status_code=404, ok=False)
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
with self.assertRaises(Exception) as context:
prompts.get("nonexistent-prompt")
self.assertIn('Prompt "nonexistent-prompt" not found', str(context.exception))
@patch("posthog.ai.prompts._get_session")
def test_handle_403_response(self, mock_get_session):
"""Should handle 403 response."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(status_code=403, ok=False)
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
with self.assertRaises(Exception) as context:
prompts.get("restricted-prompt")
self.assertIn(
'Access denied for prompt "restricted-prompt"', str(context.exception)
)
def test_throw_when_no_personal_api_key_configured(self):
"""Should throw when no personal_api_key is configured."""
posthog = self.create_mock_posthog(personal_api_key=None)
prompts = Prompts(posthog)
with self.assertRaises(Exception) as context:
prompts.get("test-prompt")
self.assertIn(
"personal_api_key is required to fetch prompts", str(context.exception)
)
@patch("posthog.ai.prompts._get_session")
def test_throw_when_api_returns_invalid_response_format(self, mock_get_session):
"""Should throw when API returns invalid response format."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(json_data={"invalid": "response"})
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
with self.assertRaises(Exception) as context:
prompts.get("test-prompt")
self.assertIn("Invalid response format", str(context.exception))
@patch("posthog.ai.prompts._get_session")
def test_use_custom_host_from_posthog_options(self, mock_get_session):
"""Should use custom host from PostHog options."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(json_data=self.mock_prompt_response)
posthog = self.create_mock_posthog(host="https://eu.i.posthog.com")
prompts = Prompts(posthog)
prompts.get("test-prompt")
call_args = mock_get.call_args
self.assertTrue(
call_args[0][0].startswith("https://eu.i.posthog.com/"),
f"Expected URL to start with 'https://eu.i.posthog.com/', got {call_args[0][0]}",
)
@patch("posthog.ai.prompts._get_session")
@patch("posthog.ai.prompts.time.time")
def test_use_default_cache_ttl_5_minutes(self, mock_time, mock_get_session):
"""Should use default cache TTL (5 minutes) when not specified."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(json_data=self.mock_prompt_response)
mock_time.return_value = 1000.0
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
# First call
prompts.get("test-prompt")
self.assertEqual(mock_get.call_count, 1)
# Advance time by 4 minutes (within default 5-minute TTL)
mock_time.return_value = 1000.0 + (4 * 60)
# Second call - should use cache
prompts.get("test-prompt")
self.assertEqual(mock_get.call_count, 1)
# Advance time past 5-minute TTL
mock_time.return_value = 1000.0 + (6 * 60)
# Third call - should refetch
prompts.get("test-prompt")
self.assertEqual(mock_get.call_count, 2)
@patch("posthog.ai.prompts._get_session")
@patch("posthog.ai.prompts.time.time")
def test_use_custom_default_cache_ttl_from_constructor(
self, mock_time, mock_get_session
):
"""Should use custom default cache TTL from constructor."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(json_data=self.mock_prompt_response)
mock_time.return_value = 1000.0
posthog = self.create_mock_posthog()
prompts = Prompts(posthog, default_cache_ttl_seconds=60)
# First call
prompts.get("test-prompt")
self.assertEqual(mock_get.call_count, 1)
# Advance time past custom TTL
mock_time.return_value = 1061.0
# Second call - should refetch
prompts.get("test-prompt")
self.assertEqual(mock_get.call_count, 2)
@patch("posthog.ai.prompts._get_session")
def test_url_encode_prompt_names_with_special_characters(self, mock_get_session):
"""Should URL-encode prompt names with special characters."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(json_data=self.mock_prompt_response)
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
prompts.get("prompt with spaces/and/slashes")
call_args = mock_get.call_args
self.assertEqual(
call_args[0][0],
"https://us.posthog.com/api/environments/@current/llm_prompts/name/prompt%20with%20spaces%2Fand%2Fslashes/",
)
@patch("posthog.ai.prompts._get_session")
def test_work_with_direct_options_no_posthog_client(self, mock_get_session):
"""Should work with direct options (no PostHog client)."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(json_data=self.mock_prompt_response)
prompts = Prompts(personal_api_key="phx_direct_key")
result = prompts.get("test-prompt")
self.assertEqual(result, self.mock_prompt_response["prompt"])
call_args = mock_get.call_args
self.assertEqual(
call_args[0][0],
"https://us.posthog.com/api/environments/@current/llm_prompts/name/test-prompt/",
)
self.assertEqual(
call_args[1]["headers"]["Authorization"], "Bearer phx_direct_key"
)
@patch("posthog.ai.prompts._get_session")
def test_use_custom_host_from_direct_options(self, mock_get_session):
"""Should use custom host from direct options."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(json_data=self.mock_prompt_response)
prompts = Prompts(
personal_api_key="phx_direct_key", host="https://eu.posthog.com"
)
prompts.get("test-prompt")
call_args = mock_get.call_args
self.assertEqual(
call_args[0][0],
"https://eu.posthog.com/api/environments/@current/llm_prompts/name/test-prompt/",
)
@patch("posthog.ai.prompts._get_session")
@patch("posthog.ai.prompts.time.time")
def test_use_custom_default_cache_ttl_from_direct_options(
self, mock_time, mock_get_session
):
"""Should use custom default cache TTL from direct options."""
mock_get = mock_get_session.return_value.get
mock_get.return_value = MockResponse(json_data=self.mock_prompt_response)
mock_time.return_value = 1000.0
prompts = Prompts(
personal_api_key="phx_direct_key", default_cache_ttl_seconds=60
)
# First call
prompts.get("test-prompt")
self.assertEqual(mock_get.call_count, 1)
# Advance time past custom TTL
mock_time.return_value = 1061.0
# Second call - should refetch
prompts.get("test-prompt")
self.assertEqual(mock_get.call_count, 2)
class TestPromptsCompile(TestPrompts):
"""Tests for the Prompts.compile() method."""
def test_replace_a_single_variable(self):
"""Should replace a single variable."""
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
result = prompts.compile("Hello, {{name}}!", {"name": "World"})
self.assertEqual(result, "Hello, World!")
def test_replace_multiple_variables(self):
"""Should replace multiple variables."""
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
result = prompts.compile(
"Hello, {{name}}! Welcome to {{company}}. Your tier is {{tier}}.",
{"name": "John", "company": "Acme Corp", "tier": "premium"},
)
self.assertEqual(
result, "Hello, John! Welcome to Acme Corp. Your tier is premium."
)
def test_handle_numbers(self):
"""Should handle numbers."""
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
result = prompts.compile("You have {{count}} items.", {"count": 42})
self.assertEqual(result, "You have 42 items.")
def test_handle_booleans(self):
"""Should handle booleans."""
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
result = prompts.compile("Feature enabled: {{enabled}}", {"enabled": True})
self.assertEqual(result, "Feature enabled: True")
def test_leave_unmatched_variables_unchanged(self):
"""Should leave unmatched variables unchanged."""
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
result = prompts.compile(
"Hello, {{name}}! Your {{unknown}} is ready.", {"name": "World"}
)
self.assertEqual(result, "Hello, World! Your {{unknown}} is ready.")
def test_handle_prompts_with_no_variables(self):
"""Should handle prompts with no variables."""
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
result = prompts.compile("You are a helpful assistant.", {})
self.assertEqual(result, "You are a helpful assistant.")
def test_handle_empty_variables_dict(self):
"""Should handle empty variables dict."""
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
result = prompts.compile("Hello, {{name}}!", {})
self.assertEqual(result, "Hello, {{name}}!")
def test_handle_multiple_occurrences_of_same_variable(self):
"""Should handle multiple occurrences of the same variable."""
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
result = prompts.compile(
"Hello, {{name}}! Goodbye, {{name}}!", {"name": "World"}
)
self.assertEqual(result, "Hello, World! Goodbye, World!")
def test_work_with_direct_options_initialization(self):
"""Should work with direct options initialization."""
prompts = Prompts(personal_api_key="phx_test_key")
result = prompts.compile("Hello, {{name}}!", {"name": "World"})
self.assertEqual(result, "Hello, World!")
def test_handle_variables_with_hyphens(self):
"""Should handle variables with hyphens."""
prompts = Prompts(personal_api_key="phx_test_key")
result = prompts.compile("User ID: {{user-id}}", {"user-id": "12345"})
self.assertEqual(result, "User ID: 12345")
def test_handle_variables_with_dots(self):
"""Should handle variables with dots."""
prompts = Prompts(personal_api_key="phx_test_key")
result = prompts.compile("Company: {{company.name}}", {"company.name": "Acme"})
self.assertEqual(result, "Company: Acme")
class TestPromptsClearCache(TestPrompts):
"""Tests for the Prompts.clear_cache() method."""
@patch("posthog.ai.prompts._get_session")
def test_clear_a_specific_prompt_from_cache(self, mock_get_session):
"""Should clear a specific prompt from cache."""
mock_get = mock_get_session.return_value.get
other_prompt_response = {**self.mock_prompt_response, "name": "other-prompt"}
mock_get.side_effect = [
MockResponse(json_data=self.mock_prompt_response),
MockResponse(json_data=other_prompt_response),
MockResponse(json_data=self.mock_prompt_response),
]
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
# Populate cache with two prompts
prompts.get("test-prompt")
prompts.get("other-prompt")
self.assertEqual(mock_get.call_count, 2)
# Clear only test-prompt
prompts.clear_cache("test-prompt")
# test-prompt should be refetched
prompts.get("test-prompt")
self.assertEqual(mock_get.call_count, 3)
# other-prompt should still be cached
prompts.get("other-prompt")
self.assertEqual(mock_get.call_count, 3)
@patch("posthog.ai.prompts._get_session")
def test_clear_all_prompts_from_cache(self, mock_get_session):
"""Should clear all prompts from cache when no name is provided."""
mock_get = mock_get_session.return_value.get
other_prompt_response = {**self.mock_prompt_response, "name": "other-prompt"}
mock_get.side_effect = [
MockResponse(json_data=self.mock_prompt_response),
MockResponse(json_data=other_prompt_response),
MockResponse(json_data=self.mock_prompt_response),
MockResponse(json_data=other_prompt_response),
]
posthog = self.create_mock_posthog()
prompts = Prompts(posthog)
# Populate cache with two prompts
prompts.get("test-prompt")
prompts.get("other-prompt")
self.assertEqual(mock_get.call_count, 2)
# Clear all cache
prompts.clear_cache()
# Both prompts should be refetched
prompts.get("test-prompt")
prompts.get("other-prompt")
self.assertEqual(mock_get.call_count, 4)
if __name__ == "__main__":
unittest.main()
+1 -1
View File
@@ -1,4 +1,4 @@
VERSION = "7.7.0"
VERSION = "7.8.2"
if __name__ == "__main__":
print(VERSION, end="") # noqa: T201
+22
View File
@@ -0,0 +1,22 @@
FROM python:3.12-slim
WORKDIR /app
# Copy the SDK source code
COPY posthog/ /app/sdk/posthog/
COPY setup.py pyproject.toml README.md LICENSE /app/sdk/
# Install the SDK from source
RUN cd /app/sdk && pip install --no-cache-dir -e .
# Install adapter dependencies
RUN pip install --no-cache-dir flask python-dateutil
# Copy adapter code
COPY sdk_compliance_adapter/adapter.py /app/adapter.py
# Expose port 8080
EXPOSE 8080
# Run the adapter
CMD ["python", "/app/adapter.py"]
+76
View File
@@ -0,0 +1,76 @@
# PostHog Python SDK Test Adapter
This adapter wraps the posthog-python SDK for compliance testing with the [PostHog SDK Test Harness](https://github.com/PostHog/posthog-sdk-test-harness).
## What is This?
This is a simple Flask app that:
1. Wraps the posthog-python SDK
2. Exposes a REST API for the test harness to control
3. Tracks internal SDK state for test assertions
## Running Tests
Tests run automatically in CI via GitHub Actions. See the test harness repo for details.
### Locally with Docker Compose
```bash
# From the posthog-python/sdk_compliance_adapter directory
docker-compose up --build --abort-on-container-exit
```
This will:
1. Build the Python SDK adapter
2. Pull the test harness image
3. Run all compliance tests
4. Show results
### Manually with Docker
```bash
# Create network
docker network create test-network
# Build and run adapter
docker build -f sdk_compliance_adapter/Dockerfile -t posthog-python-adapter .
docker run -d --name sdk-adapter --network test-network -p 8080:8080 posthog-python-adapter
# Run test harness
docker run --rm \
--name test-harness \
--network test-network \
ghcr.io/posthog/sdk-test-harness:latest \
run --adapter-url http://sdk-adapter:8080 --mock-url http://test-harness:8081
# Cleanup
docker stop sdk-adapter && docker rm sdk-adapter
docker network rm test-network
```
## Adapter Implementation
See [adapter.py](adapter.py) for the implementation.
The adapter implements the standard SDK adapter interface defined in the [test harness CONTRACT](https://github.com/PostHog/posthog-sdk-test-harness/blob/main/CONTRACT.yaml):
- `GET /health` - Return SDK information
- `POST /init` - Initialize SDK with config
- `POST /capture` - Capture an event
- `POST /flush` - Flush pending events
- `GET /state` - Return internal state
- `POST /reset` - Reset SDK state
### Key Implementation Details
**Request Tracking**: The adapter monkey-patches `batch_post` to track all HTTP requests made by the SDK, including retries.
**State Management**: Thread-safe state tracking for events captured vs sent, retry attempts, and errors.
**UUID Tracking**: Extracts and tracks UUIDs from batches to verify deduplication.
## Documentation
For complete documentation on the test harness and how to implement adapters, see:
- [PostHog SDK Test Harness](https://github.com/PostHog/posthog-sdk-test-harness)
- [Adapter Implementation Guide](https://github.com/PostHog/posthog-sdk-test-harness/blob/main/ADAPTER_GUIDE.md)
+383
View File
@@ -0,0 +1,383 @@
"""
PostHog Python SDK Test Adapter
This adapter implements the SDK Test Adapter Interface defined in the PostHog Capture API Contract.
It wraps the posthog-python SDK and exposes a REST API for the test harness to exercise.
"""
import logging
import os
import threading
import time
from typing import Any, Dict, List, Optional
from flask import Flask, jsonify, request
from posthog import Client
from posthog.request import batch_post as original_batch_post
from posthog.version import VERSION
# Configure logging
logging.basicConfig(
level=logging.DEBUG, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
app = Flask(__name__)
class RequestInfo:
"""Information about an HTTP request made by the SDK"""
def __init__(
self,
timestamp_ms: int,
status_code: int,
retry_attempt: int,
event_count: int,
uuid_list: List[str],
):
self.timestamp_ms = timestamp_ms
self.status_code = status_code
self.retry_attempt = retry_attempt
self.event_count = event_count
self.uuid_list = uuid_list
def to_dict(self) -> Dict[str, Any]:
return {
"timestamp_ms": self.timestamp_ms,
"status_code": self.status_code,
"retry_attempt": self.retry_attempt,
"event_count": self.event_count,
"uuid_list": self.uuid_list,
}
class SDKState:
"""Tracks SDK internal state for test assertions"""
def __init__(self):
self.lock = threading.Lock()
self.pending_events = 0
self.total_events_captured = 0
self.total_events_sent = 0
self.total_retries = 0
self.last_error: Optional[str] = None
self.requests_made: List[RequestInfo] = []
self.client: Optional[Client] = None
self.retry_attempts: Dict[str, int] = {} # Track retry attempts by batch ID
def reset(self):
"""Reset all state"""
with self.lock:
self.pending_events = 0
self.total_events_captured = 0
self.total_events_sent = 0
self.total_retries = 0
self.last_error = None
self.requests_made = []
self.retry_attempts = {}
if self.client:
# Flush and shutdown existing client
try:
self.client.shutdown()
except Exception as e:
logger.warning(f"Error shutting down client: {e}")
self.client = None
def increment_captured(self):
"""Increment total events captured"""
with self.lock:
self.total_events_captured += 1
self.pending_events += 1
def record_request(self, status_code: int, batch: List[Dict], batch_id: str):
"""Record an HTTP request made by the SDK"""
with self.lock:
# Determine retry attempt for this batch
retry_attempt = self.retry_attempts.get(batch_id, 0)
# Extract UUIDs from batch
uuid_list = [event.get("uuid", "") for event in batch]
request_info = RequestInfo(
timestamp_ms=int(time.time() * 1000),
status_code=status_code,
retry_attempt=retry_attempt,
event_count=len(batch),
uuid_list=uuid_list,
)
self.requests_made.append(request_info)
# Update counters
if status_code == 200:
# Success - clear pending events
self.total_events_sent += len(batch)
self.pending_events = max(0, self.pending_events - len(batch))
# Remove batch from retry tracking
self.retry_attempts.pop(batch_id, None)
else:
# Failure - increment retry count
self.retry_attempts[batch_id] = retry_attempt + 1
if retry_attempt > 0:
self.total_retries += 1
def record_error(self, error: str):
"""Record an error"""
with self.lock:
self.last_error = error
def get_state(self) -> Dict[str, Any]:
"""Get current state as dict"""
with self.lock:
return {
"pending_events": self.pending_events,
"total_events_captured": self.total_events_captured,
"total_events_sent": self.total_events_sent,
"total_retries": self.total_retries,
"last_error": self.last_error,
"requests_made": [r.to_dict() for r in self.requests_made],
}
# Global state
state = SDKState()
def create_batch_id(batch: List[Dict]) -> str:
"""Create a unique ID for a batch based on UUIDs"""
uuids = sorted([event.get("uuid", "") for event in batch])
return "-".join(uuids[:3]) # Use first 3 UUIDs as batch ID
def patched_batch_post(
api_key: str,
host: Optional[str] = None,
gzip: bool = False,
timeout: int = 15,
**kwargs,
):
"""Patched version of batch_post that tracks requests"""
batch = kwargs.get("batch", [])
batch_id = create_batch_id(batch)
try:
# Call original batch_post
response = original_batch_post(api_key, host, gzip, timeout, **kwargs)
# Record successful request
state.record_request(200, batch, batch_id)
return response
except Exception as e:
# Record failed request
status_code = (
getattr(e, "status_code", 500) if hasattr(e, "status_code") else 500
)
state.record_request(status_code, batch, batch_id)
state.record_error(str(e))
raise
# Monkey-patch the batch_post function
import posthog.request # noqa: E402
posthog.request.batch_post = patched_batch_post
# Also patch in consumer module
import posthog.consumer # noqa: E402
posthog.consumer.batch_post = patched_batch_post
@app.route("/health", methods=["GET"])
def health():
"""Health check endpoint"""
return jsonify(
{
"sdk_name": "posthog-python",
"sdk_version": VERSION,
"adapter_version": "1.0.0",
}
)
@app.route("/init", methods=["POST"])
def init():
"""Initialize the SDK client"""
try:
data = request.json or {}
# Reset state
state.reset()
# Extract config
api_key = data.get("api_key")
host = data.get("host")
flush_at = data.get("flush_at", 100)
flush_interval_ms = data.get("flush_interval_ms", 500)
max_retries = data.get("max_retries", 3)
enable_compression = data.get("enable_compression", False)
if not api_key:
return jsonify({"error": "api_key is required"}), 400
if not host:
return jsonify({"error": "host is required"}), 400
# Convert flush_interval from ms to seconds
flush_interval = flush_interval_ms / 1000.0
# Create client
client = Client(
project_api_key=api_key,
host=host,
flush_at=flush_at,
flush_interval=flush_interval,
gzip=enable_compression,
max_retries=max_retries,
debug=True,
)
state.client = client
logger.info(
f"Initialized SDK with api_key={api_key[:10]}..., host={host}, "
f"flush_at={flush_at}, flush_interval={flush_interval}, "
f"max_retries={max_retries}, gzip={enable_compression}"
)
return jsonify({"success": True})
except Exception as e:
logger.exception("Error initializing SDK")
return jsonify({"error": str(e)}), 500
@app.route("/capture", methods=["POST"])
def capture():
"""Capture a single event"""
try:
if not state.client:
return jsonify({"error": "SDK not initialized"}), 400
data = request.json or {}
distinct_id = data.get("distinct_id")
event = data.get("event")
properties = data.get("properties")
timestamp = data.get("timestamp")
if not distinct_id:
return jsonify({"error": "distinct_id is required"}), 400
if not event:
return jsonify({"error": "event is required"}), 400
# Capture event
kwargs = {"distinct_id": distinct_id, "properties": properties}
if timestamp:
# Parse ISO8601 timestamp
from dateutil.parser import parse # type: ignore[import-untyped]
kwargs["timestamp"] = parse(timestamp)
uuid = state.client.capture(event, **kwargs)
# Track that we captured an event
state.increment_captured()
logger.info(f"Captured event: {event} for {distinct_id}, uuid={uuid}")
return jsonify({"success": True, "uuid": uuid})
except Exception as e:
logger.exception("Error capturing event")
state.record_error(str(e))
return jsonify({"error": str(e)}), 500
@app.route("/identify", methods=["POST"])
def identify():
"""Identify a user"""
try:
if not state.client:
return jsonify({"error": "SDK not initialized"}), 400
data = request.json or {}
distinct_id = data.get("distinct_id")
properties = data.get("properties")
properties_set_once = data.get("properties_set_once")
if not distinct_id:
return jsonify({"error": "distinct_id is required"}), 400
# Use the identify pattern - set + set_once
if properties:
state.client.set(distinct_id=distinct_id, properties=properties)
state.increment_captured()
if properties_set_once:
state.client.set_once(
distinct_id=distinct_id, properties=properties_set_once
)
state.increment_captured()
logger.info(f"Identified user: {distinct_id}")
return jsonify({"success": True})
except Exception as e:
logger.exception("Error identifying user")
state.record_error(str(e))
return jsonify({"error": str(e)}), 500
@app.route("/flush", methods=["POST"])
def flush():
"""Force flush all pending events"""
try:
if not state.client:
return jsonify({"error": "SDK not initialized"}), 400
# Flush and wait
state.client.flush()
# Wait a bit for flush to complete
# The flush() method triggers queue.join() which blocks until all items are processed
time.sleep(0.5)
logger.info("Flushed pending events")
return jsonify({"success": True, "events_flushed": state.total_events_sent})
except Exception as e:
logger.exception("Error flushing events")
state.record_error(str(e))
return jsonify({"error": str(e), "errors": [str(e)]}, 500)
@app.route("/state", methods=["GET"])
def get_state():
"""Get internal SDK state"""
try:
return jsonify(state.get_state())
except Exception as e:
logger.exception("Error getting state")
return jsonify({"error": str(e)}), 500
@app.route("/reset", methods=["POST"])
def reset():
"""Reset SDK state"""
try:
state.reset()
logger.info("Reset SDK state")
return jsonify({"success": True})
except Exception as e:
logger.exception("Error resetting state")
return jsonify({"error": str(e)}), 500
def main():
"""Main entry point"""
port = int(os.environ.get("PORT", 8080))
logger.info(f"Starting SDK Test Adapter on port {port}")
app.run(host="0.0.0.0", port=port, debug=False)
if __name__ == "__main__":
main()
+25
View File
@@ -0,0 +1,25 @@
version: "3.8"
services:
# PostHog Python SDK adapter
sdk-adapter:
build:
context: ..
dockerfile: sdk_compliance_adapter/Dockerfile
ports:
- "8080:8080"
networks:
- test-network
# Test harness
test-harness:
image: ghcr.io/posthog/sdk-test-harness:latest
command: ["run", "--adapter-url", "http://sdk-adapter:8080", "--mock-url", "http://test-harness:8081"]
networks:
- test-network
depends_on:
- sdk-adapter
networks:
test-network:
driver: bridge
+3
View File
@@ -0,0 +1,3 @@
# SDK Test Adapter dependencies
flask>=3.0.0
python-dateutil>=2.8.0