Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions .sampo/changesets/prompts-get-all-unlabeled.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
pypi/posthog: minor
---

`prompts.get_all()` now works without a label. It fetches the latest version of every prompt in one request and warms the cache for plain `prompts.get(name)` calls. Previously the label was required, and passing `label=None` sent the literal string "None" as the label filter, returning an empty result.
66 changes: 45 additions & 21 deletions posthog/ai/prompts.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,13 @@ def _prompt_reference(
return reference


def _prompt_list_reference(label: Optional[str]) -> str:
"""Format a batch-fetch reference for logs and errors."""
if label is not None:
return f'prompts with label "{label}"'
return "all prompts"


def _extract_config(data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""Read config from an API response, tolerating servers that don't send it."""
config = data.get("config")
Expand Down Expand Up @@ -201,6 +208,9 @@ class Prompts:
# Fetch all prompts at a label in one request and warm the cache
prod_prompts = prompts.get_all(label='production')

# Fetch the latest version of every prompt in one request
all_prompts = prompts.get_all()

# Compile with variables
system_prompt = prompts.compile(template, {
'company': 'Acme Corp',
Expand Down Expand Up @@ -361,50 +371,58 @@ def get(
return fallback
raise

def get_all(self, *, label: str) -> Dict[str, PromptResult]:
def get_all(self, *, label: Optional[str] = None) -> Dict[str, PromptResult]:
"""
Fetch every prompt that carries a label, in one batch.
Fetch every prompt in one batch.

Returns a dict mapping prompt name to :class:`PromptResult`, with each
prompt at the version the label points to. Prompts without the label
are not included.
Returns a dict mapping prompt name to :class:`PromptResult`. With a
label, each prompt is at the version the label points to, and prompts
without the label are not included. Without a label, every prompt is
included at its latest version, matching what ``get(name)`` returns.

Each fetched prompt is stored in the cache, so later
``get(name, label=...)`` calls are served from cache within the TTL.
An app with many prompts can call this once per cache cycle instead of
making one ``get()`` request per prompt.
``get(name, label=...)`` (or plain ``get(name)``) calls are served
from cache within the TTL. An app with many prompts can call this once
per cache cycle instead of making one ``get()`` request per prompt.

Args:
label: The label to resolve, e.g. 'production'.
label: The label to resolve, e.g. 'production'. Omit to fetch
latest versions.

Returns:
Dict of prompt name to PromptResult.

Raises:
Exception: If the request fails, or the server does not support
fetching prompts by label on the list endpoint (PostHog
releases from before September 2026).
Exception: If the request fails, or a label was passed and the
server does not support fetching prompts by label on the list
endpoint (PostHog releases from before September 2026).
"""
try:
rows = self._fetch_prompt_list_from_api(label)
except Exception as error:
self._maybe_capture_error(error, name="*", version=None, label=label)
raise

reference = _prompt_list_reference(label)

# Validate every row before caching any, so a rejected batch leaves
# the cache untouched.
resolved_rows: List[Dict[str, Any]] = []
skipped: List[str] = []
for row in rows:
if not _is_prompt_api_response(row):
invalid_error = Exception(
f'[PostHog Prompts] Invalid response format for prompts with label "{label}"'
f"[PostHog Prompts] Invalid response format for {reference}"
)
self._maybe_capture_error(
invalid_error, name="*", version=None, label=label
)
raise invalid_error

if label is None:
resolved_rows.append(row)
continue

label_state = _row_label_state(row, label)
if label_state == "absent":
# Even one unlabeled row proves the server did not filter, and
Expand All @@ -425,7 +443,7 @@ def get_all(self, *, label: str) -> Dict[str, PromptResult]:
continue
resolved_rows.append(row)

if rows and not resolved_rows:
if label is not None and rows and not resolved_rows:
# Every returned row was skipped as moved. One moved label is a
# mid-request race, but all of them means the server most likely
# ignored the label param and served latest versions.
Expand Down Expand Up @@ -655,26 +673,32 @@ def _require_credentials(self) -> None:
"Please provide it when initializing the Prompts instance."
)

def _fetch_prompt_list_from_api(self, label: str) -> List[Dict[str, Any]]:
def _fetch_prompt_list_from_api(self, label: Optional[str]) -> List[Dict[str, Any]]:
"""
Fetch all prompts at a label from the paginated list endpoint.
Fetch all prompts from the paginated list endpoint.

Endpoint:
{host}/api/environments/@current/llm_prompts/
?token={encoded_project_api_key}&label={label}&content=full
?token={encoded_project_api_key}[&label={label}]&content=full
Auth: Bearer {personal_api_key}

Without a label the endpoint serves the latest version of every
prompt; the param must then be omitted entirely, because a literal
``label=None`` filters by a label named "None".

Follows pagination links until the last page. Returns the raw rows.
"""
self._require_credentials()

query = urllib.parse.urlencode(
{"token": self._project_api_key, "label": label, "content": "full"}
)
params: Dict[str, str] = {"token": self._project_api_key}
if label is not None:
params["label"] = label
params["content"] = "full"
query = urllib.parse.urlencode(params)
url: Optional[str] = (
f"{self._host}/api/environments/@current/llm_prompts/?{query}"
)
reference = f'prompts with label "{label}"'
reference = _prompt_list_reference(label)
headers = {
"Authorization": f"Bearer {self._personal_api_key}",
"User-Agent": USER_AGENT,
Expand Down
35 changes: 35 additions & 0 deletions posthog/test/ai/test_prompts.py
Original file line number Diff line number Diff line change
Expand Up @@ -1426,6 +1426,41 @@ def test_fetches_all_pages_and_seeds_the_cache(self, mock_get_session):
self.assertEqual(cached.source, "cache")
self.assertEqual(cached.version, 3)

@patch("posthog.ai.prompts._get_session")
def test_unlabeled_fetch_omits_the_label_param_and_seeds_the_cache(
self, mock_get_session
):
# Without a label the param must be left off the URL entirely --
# urlencode would otherwise send the literal string "label=None" and
# the server would filter by a label named "None". Rows without any
# labels must be accepted, since no label was requested.
mock_get = mock_get_session.return_value.get
unlabeled = {**self.labeled_row("prompt-a", version=2), "all_labels": []}
mock_get.return_value = self.list_response([unlabeled])

prompts = Prompts(self.create_mock_posthog())
results = prompts.get_all()

requested_url = mock_get.call_args.args[0]
self.assertNotIn("label", requested_url)
self.assertEqual(
results,
{
"prompt-a": PromptResult(
source="api",
prompt="Prompt for prompt-a",
name="prompt-a",
version=2,
)
},
)

# Later unlabeled get() calls are cache hits, not new requests.
cached = prompts.get("prompt-a", with_metadata=True)
self.assertEqual(mock_get.call_count, 1)
self.assertEqual(cached.source, "cache")
self.assertEqual(cached.version, 2)

@patch("posthog.ai.prompts._get_session")
def test_raises_when_the_server_ignores_the_label(self, mock_get_session):
# An old server ignores ?label= and returns latest versions of every
Expand Down
2 changes: 1 addition & 1 deletion references/public_api_snapshot.txt
Original file line number Diff line number Diff line change
Expand Up @@ -1348,7 +1348,7 @@ method posthog.ai.otel.processor.PostHogSpanProcessor.shutdown() -> None
method posthog.ai.prompts.Prompts.clear_cache(name: Optional[str] = None, *, version: Optional[int] = None) -> None
method posthog.ai.prompts.Prompts.compile(prompt: str, variables: PromptVariables) -> str
method posthog.ai.prompts.Prompts.get(name: str, *, with_metadata: Optional[bool] = None, cache_ttl_seconds: Optional[int] = None, fallback: Optional[str] = None, version: Optional[int] = None, label: Optional[str] = None) -> Union[str, PromptResult]
method posthog.ai.prompts.Prompts.get_all(*, label: str) -> Dict[str, PromptResult]
method posthog.ai.prompts.Prompts.get_all(*, label: Optional[str] = None) -> Dict[str, PromptResult]
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) -> Optional[str]
Expand Down