diff --git a/posthog/_async_consumer.py b/posthog/_async_consumer.py index ce0b209c..77e1294a 100644 --- a/posthog/_async_consumer.py +++ b/posthog/_async_consumer.py @@ -101,9 +101,11 @@ def __init__( flush_at: int, flush_interval: float, retries: int, - timeout: int, + timeout: float, historical_migration: bool, capture_compression: CaptureCompression, + endpoint: str = _CAPTURE_V1_PATH, + max_msg_size: int = MAX_MSG_SIZE, ) -> None: self.queue = queue self.api_key = api_key @@ -116,6 +118,8 @@ def __init__( self.timeout = timeout self.historical_migration = historical_migration self.capture_compression = capture_compression + self.endpoint = endpoint + self.max_msg_size = max_msg_size self._carryover: Optional[tuple[dict[str, Any], int]] = None self._flush_event = asyncio.Event() @@ -176,7 +180,7 @@ async def upload(self, batch: list[dict[str, Any]]) -> None: await self.request(batch) except Exception as error: await _report_capture_failure( - self.on_error, self.log, error, batch, _CAPTURE_V1_PATH + self.on_error, self.log, error, batch, self.endpoint ) finally: for _ in batch: @@ -234,12 +238,15 @@ async def next(self) -> tuple[list[dict[str, Any]], bool]: self.queue.task_done() continue - if item_size > MAX_MSG_SIZE: + if item_size > self.max_msg_size: + # Log only name and size: AI events may carry unredacted + # multimodal payloads that must not leak into logs. self.log.error( - "Event %s (%d bytes) exceeds the %dKiB limit, dropping.", + "Event %s (%d bytes) exceeds the %dKiB limit for %s, dropping.", item.get("event"), item_size, - MAX_MSG_SIZE // 1024, + self.max_msg_size // 1024, + self.endpoint, ) self.queue.task_done() continue @@ -263,4 +270,35 @@ async def request(self, batch: list[dict[str, Any]]) -> None: timeout=self.timeout, max_retries=self.retries, historical_migration=self.historical_migration, + path=self.endpoint, ) + + +class _AsyncLane: + """One capture queue, the consumer tasks that drain it, and the endpoint they post to. + + The client owns one lane per traffic class (analytics, AI), so each gets + its own backpressure, timeout, size cap and compression. + """ + + def __init__( + self, + *, + name: str, + max_queue_size: int, + endpoint: str, + max_msg_size: int, + timeout: float, + capture_compression: CaptureCompression, + ) -> None: + self.name = name + self.queue: asyncio.Queue[Any] = asyncio.Queue(max_queue_size) + self.endpoint = endpoint + self.max_msg_size = max_msg_size + self.timeout = timeout + self.capture_compression = capture_compression + self.consumers: list[_AsyncConsumer] = [] + self.worker_tasks: list[asyncio.Task[None]] = [] + + def pending_items(self) -> int: + return int(getattr(self.queue, "_unfinished_tasks", self.queue.qsize())) diff --git a/posthog/_async_request.py b/posthog/_async_request.py index 9ddc7f7a..1e10c3cd 100644 --- a/posthog/_async_request.py +++ b/posthog/_async_request.py @@ -7,7 +7,7 @@ from urllib.parse import quote from .capture_compression import CaptureCompression -from .capture_send import _parse_retry_after, _send_v1_batch +from .capture_send import _CAPTURE_V1_PATH, _parse_retry_after, _send_v1_batch from .request import ( APIError, DatetimeSerializer, @@ -150,9 +150,10 @@ async def async_send_v1_batch( batch: list[dict[str, Any]], *, compression: CaptureCompression, - timeout: int, + timeout: float, max_retries: int, historical_migration: bool, + path: str = _CAPTURE_V1_PATH, ) -> None: """Run the existing capture-v1 submitter off-loop to preserve wire parity.""" await asyncio.to_thread( @@ -164,4 +165,5 @@ async def async_send_v1_batch( timeout=timeout, max_retries=max_retries, historical_migration=historical_migration, + path=path, ) diff --git a/posthog/async_client.py b/posthog/async_client.py index 54150c0d..32e2d7bd 100644 --- a/posthog/async_client.py +++ b/posthog/async_client.py @@ -18,6 +18,7 @@ from ._async_consumer import ( _STOP, _AsyncConsumer, + _AsyncLane, _invoke_callback, _report_capture_failure, _is_processing_event, @@ -33,6 +34,7 @@ from .args import ID_TYPES, ExceptionArg, OptionalCaptureArgs, OptionalSetArgs from .capture_compression import ( CaptureCompression, + _resolve_capture_ai_compression, _resolve_capture_compression, ) from .capture_event import ( @@ -43,7 +45,7 @@ _fill_event_defaults, _merge_groups, ) -from .capture_send import _CAPTURE_V1_PATH +from .capture_send import _CAPTURE_AI_V1_PATH, _CAPTURE_V1_PATH from .client import ( MAX_DICT_SIZE as _MAX_DICT_SIZE, _MINIMAL_FLAG_CALLED_EVENT_PROPERTIES, @@ -53,6 +55,7 @@ _add_context_session_id, _context_tag_defaults, _personless_options, + _positive_config_value, add_context_tags as _add_context_tags, get_identity_state as _get_identity_state, stringify_id as _stringify_id, @@ -81,7 +84,7 @@ mark_exception_as_captured, try_attach_code_variables_to_frames, ) -from .consumer import MAX_MSG_SIZE +from .consumer import AI_MAX_MSG_SIZE, MAX_MSG_SIZE from .feature_flag_evaluations import ( FeatureFlagEvaluations, _EvaluatedFlagRecord, @@ -107,7 +110,12 @@ class AsyncClient: ``capture()`` is a synchronous, non-blocking queue write. Use ``await capture_immediate()`` when the caller must wait for delivery. + ``capture_ai()`` and ``await capture_ai_immediate()`` do the same for AI + events, on a separate queue that posts to the AI capture endpoint. ``flush()``, ``join()``, and ``shutdown()`` are awaitable lifecycle methods. + + The ``capture_ai_*``, ``privacy_mode`` and ``enable_full_ai_capture`` + options work the same way as on :class:`posthog.Client`. """ log = logging.getLogger("posthog") @@ -142,6 +150,12 @@ def __init__( code_variables_detect_secrets=None, in_app_modules: Optional[list[str]] = None, capture_compression: Optional[Union[CaptureCompression, str]] = None, + capture_ai_compression: Optional[Union[CaptureCompression, str]] = None, + capture_ai_max_queue_size: int = 1000, + capture_ai_timeout: float = 30, + capture_ai_max_event_bytes: int = AI_MAX_MSG_SIZE, + privacy_mode: bool = False, + enable_full_ai_capture: bool = False, capture_trace_context: bool = False, secret_key: Optional[str] = None, personal_api_key: Optional[str] = None, @@ -152,6 +166,18 @@ def __init__( raise ValueError("flush_at must be greater than zero") if flush_interval <= 0: raise ValueError("flush_interval must be greater than zero") + capture_ai_max_queue_size = _positive_config_value( + "capture_ai_max_queue_size", capture_ai_max_queue_size, integer=True + ) + capture_ai_timeout = _positive_config_value( + "capture_ai_timeout", capture_ai_timeout + ) + capture_ai_max_event_bytes = _positive_config_value( + "capture_ai_max_event_bytes", + capture_ai_max_event_bytes, + integer=True, + maximum=AI_MAX_MSG_SIZE, + ) self.api_key = (project_api_key or "").strip() self.raw_host = normalize_host(host) @@ -169,7 +195,12 @@ def __init__( self.super_options = super_options self._release_id = _resolve_release_id() self.capture_compression = _resolve_capture_compression(capture_compression) + self.capture_ai_compression = _resolve_capture_ai_compression( + capture_ai_compression + ) self.capture_trace_context = capture_trace_context + self.privacy_mode = privacy_mode + self.enable_full_ai_capture = enable_full_ai_capture is True if personal_api_key is not None and secret_key is None: warnings.warn( "`personal_api_key` is deprecated; use `secret_key` instead.", @@ -244,12 +275,29 @@ def __init__( "api_key is empty after trimming whitespace; check your project API key" ) - self._queue: asyncio.Queue[Any] = asyncio.Queue(max_queue_size) + self._analytics_lane = _AsyncLane( + name="analytics", + max_queue_size=max_queue_size, + endpoint=_CAPTURE_V1_PATH, + max_msg_size=MAX_MSG_SIZE, + timeout=timeout, + capture_compression=self.capture_compression, + ) + # AI events post to their own endpoint, so multi-MB events stay off the + # analytics endpoint's smaller caps. Its workers start on the first AI + # event, so clients that never send one pay nothing. + self._ai_lane = _AsyncLane( + name="ai", + max_queue_size=capture_ai_max_queue_size, + endpoint=_CAPTURE_AI_V1_PATH, + max_msg_size=capture_ai_max_event_bytes, + timeout=capture_ai_timeout, + capture_compression=self.capture_ai_compression, + ) + self._lanes = (self._analytics_lane, self._ai_lane) self._worker_count = max(1, thread) self._flush_at = flush_at self._flush_interval = flush_interval - self._consumers: list[_AsyncConsumer] = [] - self._worker_tasks: list[asyncio.Task[None]] = [] self._immediate_callers: dict[asyncio.Task[Any], int] = {} self._inflight_operations: set[asyncio.Future[None]] = set() self._http_client: Optional[Any] = None @@ -263,7 +311,7 @@ def __init__( self._register_duplicate_client() async def __aenter__(self) -> AsyncClient: - self._ensure_workers_started() + self._ensure_workers_started(self._analytics_lane) return self async def __aexit__(self, exc_type, exc, tb) -> bool: @@ -324,9 +372,9 @@ def _get_http_client(self): self._http_client = _build_client(self.host) return self._http_client - def _new_consumer(self) -> _AsyncConsumer: + def _new_consumer(self, lane: _AsyncLane) -> _AsyncConsumer: return _AsyncConsumer( - self._queue, + lane.queue, self.api_key, host=self.host, on_error=self.on_error, @@ -334,22 +382,27 @@ def _new_consumer(self) -> _AsyncConsumer: flush_at=self._flush_at, flush_interval=self._flush_interval, retries=self.max_retries, - timeout=self.timeout, + timeout=lane.timeout, historical_migration=self.historical_migration, - capture_compression=self.capture_compression, + capture_compression=lane.capture_compression, + endpoint=lane.endpoint, + max_msg_size=lane.max_msg_size, ) - def _ensure_workers_started(self) -> None: - if self.disabled or not self.send or self._closed or self._worker_tasks: + def _ensure_workers_started(self, lane: _AsyncLane) -> None: + if self.disabled or not self.send or self._closed or lane.worker_tasks: return self._bind_loop() for _ in range(self._worker_count): - consumer = self._new_consumer() - self._consumers.append(consumer) - self._worker_tasks.append(asyncio.create_task(consumer.run())) + consumer = self._new_consumer(lane) + lane.consumers.append(consumer) + lane.worker_tasks.append(asyncio.create_task(consumer.run())) + + def _all_worker_tasks(self) -> list[asyncio.Task[None]]: + return [task for lane in self._lanes for task in lane.worker_tasks] def _enqueue_prepared_event( - self, prepared: dict[str, Any], defaults: _EventDefaults + self, prepared: dict[str, Any], lane: _AsyncLane, defaults: _EventDefaults ) -> bool: if not self._accepting or self._closed: return False @@ -361,13 +414,13 @@ def _enqueue_prepared_event( if self._loop is None: if running_loop is not None: - self._ensure_workers_started() - self._queue.put_nowait(queued_event) + self._ensure_workers_started(lane) + lane.queue.put_nowait(queued_event) return True if running_loop is self._loop: - self._ensure_workers_started() - self._queue.put_nowait(queued_event) + self._ensure_workers_started(lane) + lane.queue.put_nowait(queued_event) return True if running_loop is not None: raise RuntimeError("AsyncClient cannot be shared across event loops") @@ -381,8 +434,8 @@ def enqueue_on_bound_loop() -> None: if not self._accepting or self._closed: admitted.set_result(False) return - self._ensure_workers_started() - self._queue.put_nowait(queued_event) + self._ensure_workers_started(lane) + lane.queue.put_nowait(queued_event) except BaseException as error: admitted.set_exception(error) else: @@ -542,6 +595,30 @@ def capture( self, event: str, **kwargs: Unpack[OptionalCaptureArgs] ) -> Optional[str]: """Queue an event without blocking for network delivery.""" + return self._capture(event, self._analytics_lane, kwargs) + + def capture_ai( + self, event: str, **kwargs: Unpack[OptionalCaptureArgs] + ) -> Optional[str]: + """Queue an AI event for the AI capture endpoint without blocking. + + Takes the same arguments and returns the same value as ``capture()``. + The event uses a separate queue with a larger per-event size cap. The + payload is sent as given, with no redaction or truncation. + """ + self._log_non_ai_event(event) + return self._capture(event, self._ai_lane, kwargs) + + def _log_non_ai_event(self, event: str) -> None: + if isinstance(event, str) and not event.startswith("$ai_"): + self.log.debug( + "capture_ai called with non-AI event name %r; routing it to the AI endpoint anyway.", + event, + ) + + def _capture( + self, event: str, lane: _AsyncLane, kwargs: OptionalCaptureArgs + ) -> Optional[str]: try: msg, property_allowlist, defaults = self._build_capture_event(event, kwargs) prepared, sent_uuid = self._prepare_event(msg, property_allowlist) @@ -550,12 +627,12 @@ def capture( if not self.send: return sent_uuid - if not self._enqueue_prepared_event(prepared, defaults): + if not self._enqueue_prepared_event(prepared, lane, defaults): return None self.log.debug("queued async event %s", event) return sent_uuid except asyncio.QueueFull: - self.log.warning("PostHog async capture queue is full") + self.log.warning("PostHog async %s capture queue is full", lane.name) return None except Exception as error: if self.debug: @@ -577,6 +654,21 @@ async def capture_immediate( self, event: str, **kwargs: Unpack[OptionalCaptureArgs] ) -> Optional[str]: """Capture one event and wait until its delivery attempt completes.""" + return await self._capture_immediate(event, self._analytics_lane, kwargs) + + async def capture_ai_immediate( + self, event: str, **kwargs: Unpack[OptionalCaptureArgs] + ) -> Optional[str]: + """Capture one AI event and wait until its delivery attempt completes. + + Uses the AI capture endpoint and the ``capture_ai_*`` settings. + """ + self._log_non_ai_event(event) + return await self._capture_immediate(event, self._ai_lane, kwargs) + + async def _capture_immediate( + self, event: str, lane: _AsyncLane, kwargs: OptionalCaptureArgs + ) -> Optional[str]: current = asyncio.current_task() if current is None: # pragma: no cover - async functions always have a task return None @@ -603,21 +695,22 @@ async def capture_immediate( "unable to serialize immediate event for sizing, dropping" ) return None - if event_size > MAX_MSG_SIZE: + if event_size > lane.max_msg_size: self.log.error( - "Event %s (%d bytes) exceeds the %dKiB limit, dropping.", + "Event %s (%d bytes) exceeds the %dKiB limit for %s, dropping.", processed.get("event"), event_size, - MAX_MSG_SIZE // 1024, + lane.max_msg_size // 1024, + lane.endpoint, ) return None - consumer = self._new_consumer() + consumer = self._new_consumer(lane) await consumer.request(error_batch) return sent_uuid except Exception as error: await _report_capture_failure( - self.on_error, self.log, error, error_batch, _CAPTURE_V1_PATH + self.on_error, self.log, error, error_batch, lane.endpoint ) if self.debug: raise @@ -773,7 +866,7 @@ def _enqueue_built_event( defaults = self._event_defaults( context_options=context_options, disable_geoip=disable_geoip ) - if not self._enqueue_prepared_event(prepared, defaults): + if not self._enqueue_prepared_event(prepared, self._analytics_lane, defaults): return None return sent_uuid @@ -1100,40 +1193,43 @@ def _capture_feature_flag_called_if_needed( reported_flags.add(reported_key) def _pending_queue_items(self) -> int: - return int(getattr(self._queue, "_unfinished_tasks", self._queue.qsize())) + return sum(lane.pending_items() for lane in self._lanes) def _defer_lifecycle_call(self, awaitable) -> None: task = asyncio.create_task(_run_outside_processing_event(awaitable)) self._deferred_lifecycle_tasks.add(task) task.add_done_callback(self._deferred_lifecycle_tasks.discard) - def _discard_undrainable_queue(self) -> None: + def _discard_undrainable_queue(self, lane: _AsyncLane) -> None: discarded = 0 while True: try: - self._queue.get_nowait() + lane.queue.get_nowait() except asyncio.QueueEmpty: break - self._queue.task_done() + lane.queue.task_done() discarded += 1 - orphaned = self._pending_queue_items() + orphaned = lane.pending_items() for _ in range(orphaned): - self._queue.task_done() + lane.queue.task_done() discarded += orphaned if discarded: self.log.warning( - "discarded %d async capture items because all workers exited", + "discarded %d async %s capture items because all workers exited", discarded, + lane.name, ) - async def _wait_for_queue_drain(self, deadline: Optional[float]) -> None: - live_workers = [task for task in self._worker_tasks if not task.done()] + async def _wait_for_queue_drain( + self, lane: _AsyncLane, deadline: Optional[float] + ) -> None: + live_workers = [task for task in lane.worker_tasks if not task.done()] if not live_workers: - self._discard_undrainable_queue() + self._discard_undrainable_queue(lane) return - queue_join = asyncio.create_task(self._queue.join()) + queue_join = asyncio.create_task(lane.queue.join()) async def wait_for_workers() -> None: await asyncio.wait(live_workers, return_when=asyncio.ALL_COMPLETED) @@ -1156,7 +1252,7 @@ async def wait_for_workers() -> None: await queue_join return - self._discard_undrainable_queue() + self._discard_undrainable_queue(lane) await queue_join finally: for task in (queue_join, workers_finished): @@ -1165,22 +1261,28 @@ async def wait_for_workers() -> None: await asyncio.gather(queue_join, workers_finished, return_exceptions=True) async def flush(self, timeout_seconds: Optional[float] = 10) -> None: - if asyncio.current_task() in self._worker_tasks or _is_processing_event(): + if asyncio.current_task() in self._all_worker_tasks() or _is_processing_event(): self._defer_lifecycle_call(self.flush(timeout_seconds)) return if not self.send or self.disabled or self._pending_queue_items() == 0: return - self._ensure_workers_started() + pending_lanes = [lane for lane in self._lanes if lane.pending_items()] + for lane in pending_lanes: + self._ensure_workers_started(lane) deadline = ( None if timeout_seconds is None else asyncio.get_running_loop().time() + timeout_seconds ) try: - for consumer in self._consumers: - consumer.request_flush() - - await self._wait_for_queue_drain(deadline) + # Wake every lane before waiting on any, so one lane's partial + # batch does not wait out its flush_interval behind the other. + for lane in pending_lanes: + for consumer in lane.consumers: + consumer.request_flush() + + for lane in pending_lanes: + await self._wait_for_queue_drain(lane, deadline) except asyncio.TimeoutError: self.log.warning( "flush timed out after %s seconds with %s items pending", @@ -1197,7 +1299,7 @@ async def _close_transport(self) -> None: async def shutdown(self) -> None: current = asyncio.current_task() if ( - current in self._worker_tasks + current in self._all_worker_tasks() or current in self._immediate_callers or _is_processing_event() ): @@ -1227,17 +1329,20 @@ async def shutdown(self) -> None: errors.append(error) try: - live_workers = [task for task in self._worker_tasks if not task.done()] - for _ in live_workers: - await self._queue.put(_STOP) - if self._worker_tasks: - await asyncio.gather(*self._worker_tasks, return_exceptions=True) + for lane in self._lanes: + live_workers = [t for t in lane.worker_tasks if not t.done()] + for _ in live_workers: + await lane.queue.put(_STOP) + worker_tasks = self._all_worker_tasks() + if worker_tasks: + await asyncio.gather(*worker_tasks, return_exceptions=True) except Exception as error: self.log.exception("Failed to stop async capture workers") errors.append(error) finally: - self._worker_tasks.clear() - self._consumers.clear() + for lane in self._lanes: + lane.worker_tasks.clear() + lane.consumers.clear() try: await self._close_transport() diff --git a/posthog/capture_send.py b/posthog/capture_send.py index f254f165..3c5f916d 100644 --- a/posthog/capture_send.py +++ b/posthog/capture_send.py @@ -283,7 +283,7 @@ def _post_v1( attempt: int, request_id: str, compression: CaptureCompression = CaptureCompression.NONE, - timeout: int = 15, + timeout: float = 15, sdk_info: str = USER_AGENT, session: Optional["requests.Session"] = None, path: str = _CAPTURE_V1_PATH, @@ -405,7 +405,7 @@ def _send_v1_batch( batch: list[dict], *, compression: CaptureCompression = CaptureCompression.NONE, - timeout: int = 15, + timeout: float = 15, max_retries: int = 3, historical_migration: bool = False, sdk_info: str = USER_AGENT, diff --git a/posthog/test/test_ai_capture_lane.py b/posthog/test/test_ai_capture_lane.py index dfeedcb8..885d4440 100644 --- a/posthog/test/test_ai_capture_lane.py +++ b/posthog/test/test_ai_capture_lane.py @@ -9,6 +9,7 @@ import posthog from posthog.ai.utils import _capture_ai_event, finalize_ai_content, with_privacy_mode +from posthog.async_client import AsyncPosthog from posthog.capture_compression import CAPTURE_COMPRESSION_ENV_VAR, CaptureCompression from posthog.client import Client from posthog.consumer import AI_MAX_MSG_SIZE, AI_MAX_PROPERTIES_SIZE, MAX_MSG_SIZE @@ -697,14 +698,16 @@ def drop_uuid(event): class TestCaptureAiPrivacyMode(unittest.TestCase): """Privacy mode always wins over `enable_full_ai_capture`.""" - def test_privacy_mode_strips_content_despite_full_ai_capture(self): - client = Client( + @parameterized.expand([("sync", Client), ("async", AsyncPosthog)]) + def test_privacy_mode_strips_content_despite_full_ai_capture( + self, _name, client_cls + ): + client = client_cls( TEST_API_KEY, send=False, enable_full_ai_capture=True, privacy_mode=True, ) - self.addCleanup(client.join) payload = {"role": "user", "content": "sensitive prompt"} sanitized = with_privacy_mode( @@ -713,6 +716,33 @@ def test_privacy_mode_strips_content_despite_full_ai_capture(self): self.assertIsNone(sanitized) + @parameterized.expand( + [ + ("sync_default", Client, False), + ("sync_full_capture", Client, True), + ("async_default", AsyncPosthog, False), + ("async_full_capture", AsyncPosthog, True), + ] + ) + def test_full_ai_capture_controls_media_redaction( + self, _name, client_cls, full_capture + ): + client = client_cls( + TEST_API_KEY, send=False, enable_full_ai_capture=full_capture + ) + image = "data:image/jpeg;base64," + "A" * 64 + payload = [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": image}}], + } + ] + + sanitized = finalize_ai_content(payload, ph_client=client) + + sent_url = sanitized[0]["content"][0]["image_url"]["url"] + self.assertEqual(sent_url == image, full_capture) + if __name__ == "__main__": unittest.main() diff --git a/posthog/test/test_async_client.py b/posthog/test/test_async_client.py index 02ce2ce4..27c370a2 100644 --- a/posthog/test/test_async_client.py +++ b/posthog/test/test_async_client.py @@ -12,7 +12,8 @@ import pytest from posthog import AsyncClient, AsyncPosthog, CaptureCompression -from posthog.consumer import MAX_MSG_SIZE +from posthog.capture_send import _CAPTURE_AI_V1_PATH, _CAPTURE_V1_PATH +from posthog.consumer import AI_MAX_MSG_SIZE, MAX_MSG_SIZE from posthog.contexts import ( new_context, set_capture_exception_code_variables_context, @@ -88,8 +89,8 @@ async def send_batch(api_key, host, batch, **kwargs): with patch_async_capture_send(side_effect=send_batch): client = AsyncPosthog("test-key", flush_interval=30) - client._ensure_workers_started() - while not client._queue._getters: + client._ensure_workers_started(client._analytics_lane) + while not client._analytics_lane.queue._getters: await asyncio.sleep(0) asyncio.get_running_loop().set_debug(True) @@ -124,7 +125,7 @@ async def send_batch(api_key, host, batch, **kwargs): @pytest.mark.asyncio async def test_cross_thread_capture_rechecks_shutdown_before_queue_admission(): client = AsyncPosthog("test-key", flush_interval=30) - client._ensure_workers_started() + client._ensure_workers_started(client._analytics_lane) scheduled_callbacks = [] capture_result = [] @@ -148,9 +149,9 @@ async def test_cross_thread_capture_rechecks_shutdown_before_queue_admission(): capture_thread.join(timeout=1) assert capture_result == [None] - for task in client._worker_tasks: + for task in client._analytics_lane.worker_tasks: task.cancel() - await asyncio.gather(*client._worker_tasks, return_exceptions=True) + await asyncio.gather(*client._analytics_lane.worker_tasks, return_exceptions=True) await client._close_transport() @@ -283,22 +284,41 @@ async def send_batch(api_key, host, batch, **kwargs): @pytest.mark.asyncio -async def test_capture_immediate_drops_oversized_event_after_before_send(): +@pytest.mark.parametrize( + ("method_name", "client_kwargs", "payload_size", "sent"), + [ + ("capture_immediate", {}, MAX_MSG_SIZE, False), + ("capture_ai_immediate", {}, MAX_MSG_SIZE, True), + ("capture_ai_immediate", {"capture_ai_max_event_bytes": 1024}, 2048, False), + ], +) +async def test_capture_immediate_applies_its_lane_size_cap_after_before_send( + method_name, client_kwargs, payload_size, sent +): def before_send(event): - event["properties"]["user_input"] = "x" * MAX_MSG_SIZE + event["properties"]["user_input"] = "x" * payload_size return event with patch_async_capture_send(new=mock.AsyncMock()) as send_batch: - client = AsyncPosthog("test-key", before_send=before_send) - result = await client.capture_immediate("event", distinct_id="user-1") + client = AsyncPosthog("test-key", before_send=before_send, **client_kwargs) + result = await getattr(client, method_name)("event", distinct_id="user-1") await client.shutdown() - assert result is None - send_batch.assert_not_awaited() + assert (result is not None) is sent + assert send_batch.await_count == (1 if sent else 0) @pytest.mark.asyncio -async def test_capture_immediate_uses_capture_v1_without_building_httpx_client(): +@pytest.mark.parametrize( + ("method_name", "path", "compression", "timeout"), + [ + ("capture_immediate", _CAPTURE_V1_PATH, CaptureCompression.GZIP, 15), + ("capture_ai_immediate", _CAPTURE_AI_V1_PATH, CaptureCompression.DEFLATE, 45), + ], +) +async def test_capture_immediate_uses_its_lane_without_building_httpx_client( + method_name, path, compression, timeout +): with ( mock.patch( "posthog._async_consumer.async_send_v1_batch", new=mock.AsyncMock() @@ -308,24 +328,72 @@ async def test_capture_immediate_uses_capture_v1_without_building_httpx_client() client = AsyncPosthog( "test-key", capture_compression=CaptureCompression.GZIP, + capture_ai_compression=CaptureCompression.DEFLATE, + capture_ai_timeout=45, + ) + event_uuid = await getattr(client, method_name)( + "$ai_generation", distinct_id="user-1" ) - event_uuid = await client.capture_immediate("event", distinct_id="user-1") await client.shutdown() assert event_uuid is not None build_client.assert_not_called() send_v1.assert_awaited_once() - assert send_v1.await_args.kwargs["compression"] == CaptureCompression.GZIP + assert send_v1.await_args.kwargs["path"] == path + assert send_v1.await_args.kwargs["compression"] == compression + assert send_v1.await_args.kwargs["timeout"] == timeout assert send_v1.await_args.args[2][0]["uuid"] == event_uuid +@pytest.mark.asyncio +async def test_ai_events_queue_on_their_own_lane_and_flush_drains_both(): + sent: dict[str, list[str]] = {} + + async def send_batch(api_key, host, batch, **kwargs): + sent.setdefault(kwargs["path"], []).extend(e["event"] for e in batch) + + with patch_async_capture_send(side_effect=send_batch): + client = AsyncPosthog("test-key", flush_at=100, flush_interval=30) + client.capture("pageview", distinct_id="user-1") + assert client._ai_lane.worker_tasks == [] + client.capture_ai( + "$ai_generation", + distinct_id="user-1", + properties={"$ai_input": "x" * MAX_MSG_SIZE}, + ) + await client.flush(timeout_seconds=1) + assert sent == { + _CAPTURE_V1_PATH: ["pageview"], + _CAPTURE_AI_V1_PATH: ["$ai_generation"], + } + + client.capture_ai("$ai_span", distinct_id="user-1") + await client.shutdown() + + assert sent[_CAPTURE_AI_V1_PATH] == ["$ai_generation", "$ai_span"] + assert client._all_worker_tasks() == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("event", [None, 123]) +async def test_ai_capture_handles_a_non_string_event_name_like_capture(event): + client = AsyncPosthog("test-key", send=False) + try: + expected = client.capture(event, distinct_id="user-1") is None + assert (client.capture_ai(event, distinct_id="user-1") is None) is expected + immediate = await client.capture_ai_immediate(event, distinct_id="user-1") + assert (immediate is None) is expected + finally: + await client.shutdown() + + @pytest.mark.asyncio async def test_send_false_accepts_without_starting_workers_or_transport(): with mock.patch("posthog.async_client._build_client") as build_client: client = AsyncPosthog("test-key", send=False) assert client.capture("event", distinct_id="user-1") is not None assert await client.capture_immediate("event", distinct_id="user-1") is not None - assert client._worker_tasks == [] + assert client._all_worker_tasks() == [] await client.shutdown() build_client.assert_not_called() @@ -380,7 +448,7 @@ async def test_capture_after_shutdown_is_dropped_without_restarting_workers(): await client.shutdown() assert client.capture("event", distinct_id="user-1") is None assert await client.capture_immediate("event", distinct_id="user-1") is None - assert client._worker_tasks == [] + assert client._all_worker_tasks() == [] build_client.assert_not_called() @@ -521,10 +589,10 @@ async def send_batch(api_key, host, batch, **kwargs): @pytest.mark.asyncio async def test_shutdown_returns_when_all_workers_exited_with_queued_work(): client = AsyncPosthog("test-key", flush_interval=30) - client._ensure_workers_started() - for task in client._worker_tasks: + client._ensure_workers_started(client._analytics_lane) + for task in client._analytics_lane.worker_tasks: task.cancel() - await asyncio.gather(*client._worker_tasks, return_exceptions=True) + await asyncio.gather(*client._analytics_lane.worker_tasks, return_exceptions=True) assert client.capture("undrainable event", distinct_id="user-1") is not None try: @@ -625,8 +693,17 @@ async def flush_and_close(): assert batches[0][0]["event"] == "event" -@pytest.mark.parametrize(("option", "value"), [("flush_at", 0), ("flush_interval", 0)]) -def test_rejects_non_positive_batch_settings(option, value): +@pytest.mark.parametrize( + ("option", "value"), + [ + ("flush_at", 0), + ("flush_interval", 0), + ("capture_ai_max_queue_size", 0), + ("capture_ai_timeout", 0), + ("capture_ai_max_event_bytes", AI_MAX_MSG_SIZE + 1), + ], +) +def test_rejects_invalid_batch_settings(option, value): with pytest.raises(ValueError, match=option): AsyncPosthog("test-key", **{option: value}) diff --git a/references/public_api_snapshot.txt b/references/public_api_snapshot.txt index a0214f6f..516b1e71 100644 --- a/references/public_api_snapshot.txt +++ b/references/public_api_snapshot.txt @@ -656,6 +656,7 @@ attribute posthog.args.OptionalSetArgs.timestamp: NotRequired[Optional[Union[dat attribute posthog.args.OptionalSetArgs.uuid: NotRequired[Optional[Union[str, UUID]]] attribute posthog.async_client.AsyncClient.api_key = (project_api_key or '').strip() attribute posthog.async_client.AsyncClient.before_send = before_send +attribute posthog.async_client.AsyncClient.capture_ai_compression = _resolve_capture_ai_compression(capture_ai_compression) attribute posthog.async_client.AsyncClient.capture_compression = _resolve_capture_compression(capture_compression) attribute posthog.async_client.AsyncClient.capture_exception_code_variables = capture_exception_code_variables attribute posthog.async_client.AsyncClient.capture_trace_context = capture_trace_context @@ -667,6 +668,7 @@ attribute posthog.async_client.AsyncClient.debug = debug attribute posthog.async_client.AsyncClient.disable_geoip = disable_geoip attribute posthog.async_client.AsyncClient.disabled = disabled or not self.api_key attribute posthog.async_client.AsyncClient.distinct_ids_feature_flags_reported = SizeLimitedDict(_MAX_DICT_SIZE, set) +attribute posthog.async_client.AsyncClient.enable_full_ai_capture = enable_full_ai_capture is True attribute posthog.async_client.AsyncClient.feature_flags_request_max_retries = max(0, feature_flags_request_max_retries) attribute posthog.async_client.AsyncClient.feature_flags_request_timeout_seconds = feature_flags_request_timeout_seconds attribute posthog.async_client.AsyncClient.historical_migration = historical_migration @@ -678,6 +680,7 @@ attribute posthog.async_client.AsyncClient.log_captured_exceptions = log_capture attribute posthog.async_client.AsyncClient.max_retries = max(0, max_retries) attribute posthog.async_client.AsyncClient.on_error = on_error attribute posthog.async_client.AsyncClient.personal_api_key = self.secret_key +attribute posthog.async_client.AsyncClient.privacy_mode = privacy_mode attribute posthog.async_client.AsyncClient.project_root = os.getcwd() attribute posthog.async_client.AsyncClient.raw_host = normalize_host(host) attribute posthog.async_client.AsyncClient.secret_key = (resolved_secret_key.strip() if isinstance(resolved_secret_key, str) else resolved_secret_key) or None @@ -1176,7 +1179,7 @@ class posthog.ai.types.TokenUsage class posthog.ai.types.ToolInProgress class posthog.args.OptionalCaptureArgs class posthog.args.OptionalSetArgs -class posthog.async_client.AsyncClient(project_api_key: str, host: Optional[str] = None, *, debug: bool = False, max_queue_size: int = 10000, send: bool = True, on_error=None, flush_at: int = 100, flush_interval: float = 5.0, max_retries: int = 3, timeout: int = 15, thread: int = 1, disabled: bool = False, disable_geoip: bool = True, is_server: bool = True, historical_migration: bool = False, super_properties: Optional[dict[str, Any]] = None, super_options: Optional[dict[str, Any]] = None, before_send=None, log_captured_exceptions: bool = False, project_root: Optional[str] = None, capture_exception_code_variables: bool = False, code_variables_mask_patterns=None, code_variables_ignore_patterns=None, code_variables_mask_url_credentials=None, code_variables_detect_secrets=None, in_app_modules: Optional[list[str]] = None, capture_compression: Optional[Union[CaptureCompression, str]] = None, capture_trace_context: bool = False, secret_key: Optional[str] = None, personal_api_key: Optional[str] = None, feature_flags_request_timeout_seconds: int = 3, feature_flags_request_max_retries: int = 1) +class posthog.async_client.AsyncClient(project_api_key: str, host: Optional[str] = None, *, debug: bool = False, max_queue_size: int = 10000, send: bool = True, on_error=None, flush_at: int = 100, flush_interval: float = 5.0, max_retries: int = 3, timeout: int = 15, thread: int = 1, disabled: bool = False, disable_geoip: bool = True, is_server: bool = True, historical_migration: bool = False, super_properties: Optional[dict[str, Any]] = None, super_options: Optional[dict[str, Any]] = None, before_send=None, log_captured_exceptions: bool = False, project_root: Optional[str] = None, capture_exception_code_variables: bool = False, code_variables_mask_patterns=None, code_variables_ignore_patterns=None, code_variables_mask_url_credentials=None, code_variables_detect_secrets=None, in_app_modules: Optional[list[str]] = None, capture_compression: Optional[Union[CaptureCompression, str]] = None, capture_ai_compression: Optional[Union[CaptureCompression, str]] = None, capture_ai_max_queue_size: int = 1000, capture_ai_timeout: float = 30, capture_ai_max_event_bytes: int = AI_MAX_MSG_SIZE, privacy_mode: bool = False, enable_full_ai_capture: bool = False, capture_trace_context: bool = False, secret_key: Optional[str] = None, personal_api_key: Optional[str] = None, feature_flags_request_timeout_seconds: int = 3, feature_flags_request_max_retries: int = 1) class posthog.async_client.AsyncPosthog class posthog.bucketed_rate_limiter.BucketedRateLimiter(bucket_size: Number, refill_rate: Number, refill_interval_seconds: Number, on_bucket_rate_limited: Optional[Callable[[Hashable], None]] = None, clock: Callable[[], float] = time.monotonic) class posthog.capture_compression.CaptureCompression @@ -1568,6 +1571,8 @@ method posthog.ai.stream.AsyncStreamWrapper.aclose() -> None method posthog.ai.stream.AsyncStreamWrapper.close() -> None method posthog.async_client.AsyncClient.alias(previous_id: ID_TYPES, distinct_id: Optional[str], timestamp: Optional[Union[datetime, str]] = None, uuid: Optional[str] = None, disable_geoip: Optional[bool] = None, options: Optional[Dict[str, Any]] = None) -> Optional[str] method posthog.async_client.AsyncClient.capture(event: str, **kwargs: Unpack[OptionalCaptureArgs]) -> Optional[str] +method posthog.async_client.AsyncClient.capture_ai(event: str, **kwargs: Unpack[OptionalCaptureArgs]) -> Optional[str] +method posthog.async_client.AsyncClient.capture_ai_immediate(event: str, **kwargs: Unpack[OptionalCaptureArgs]) -> Optional[str] method posthog.async_client.AsyncClient.capture_exception(exception: Optional[ExceptionArg] = None, **kwargs: Unpack[OptionalCaptureArgs]) -> Optional[str] method posthog.async_client.AsyncClient.capture_immediate(event: str, **kwargs: Unpack[OptionalCaptureArgs]) -> Optional[str] method posthog.async_client.AsyncClient.evaluate_flags(distinct_id: Optional[ID_TYPES] = None, *, groups: Optional[Mapping[str, Union[str, int]]] = None, person_properties: Optional[Dict[str, Any]] = None, group_properties: Optional[Dict[str, Dict[str, Any]]] = None, disable_geoip: Optional[bool] = None, flag_keys: Optional[list[str]] = None, device_id: Optional[str] = None) -> FeatureFlagEvaluations diff --git a/sdk_compliance_adapter/adapter.py b/sdk_compliance_adapter/adapter.py index 19068152..7a9ec437 100644 --- a/sdk_compliance_adapter/adapter.py +++ b/sdk_compliance_adapter/adapter.py @@ -172,7 +172,7 @@ def patched_post_v1( attempt: int, request_id: str, compression: CaptureCompression = CaptureCompression.NONE, - timeout: int = 15, + timeout: float = 15, sdk_info: str = USER_AGENT, session: Any = None, path: str = _CAPTURE_V1_PATH, diff --git a/typings/requests/__init__.pyi b/typings/requests/__init__.pyi index 9e7892be..4cc05e76 100644 --- a/typings/requests/__init__.pyi +++ b/typings/requests/__init__.pyi @@ -34,7 +34,7 @@ class Session: *, data: str | bytes, headers: dict[str, str], - timeout: int, + timeout: float, stream: bool = ..., ) -> Response: ... def get(