Merge remote-tracking branch 'origin/main' into fix/call-concurrency-limit

# Conflicts:
#	api/services/auth/depends.py
#	docs/api-reference/openapi.json
#	sdk/python/src/dograh_sdk/_generated_models.py
This commit is contained in:
Abhishek Kumar 2026-07-09 18:26:20 +05:30
commit d3326d8fad
57 changed files with 1163 additions and 1354 deletions

View file

@ -34,9 +34,6 @@ async def require_local_auth() -> None:
raise HTTPException(status_code=404, detail="Not found")
POSTHOG_ORGANIZATION_USES_MPS_BILLING_V2_PROPERTY = "uses_mps_billing_v2"
async def get_user(
authorization: Annotated[str | None, Header()] = None,
x_api_key: Annotated[str | None, Header(alias="X-API-Key")] = None,
@ -193,7 +190,6 @@ def _sync_created_organization_to_posthog(
organization,
stack_user: dict | None = None,
created_by_provider_id: str | None = None,
uses_mps_billing_v2: bool | None = None,
) -> None:
"""Create/update the PostHog organization group for a newly-created org."""
try:
@ -209,10 +205,6 @@ def _sync_created_organization_to_posthog(
}
if created_by:
properties["created_by_provider_id"] = created_by
if uses_mps_billing_v2 is not None:
properties[POSTHOG_ORGANIZATION_USES_MPS_BILLING_V2_PROPERTY] = (
uses_mps_billing_v2
)
group_identify(
POSTHOG_ORGANIZATION_GROUP_TYPE,
@ -231,50 +223,6 @@ def _sync_created_organization_to_posthog(
logger.exception("Failed to sync created organization to PostHog")
def _sync_posthog_organization_group_properties(
*,
organization,
uses_mps_billing_v2: bool | None = None,
) -> None:
"""Update PostHog organization group properties without creating a person."""
try:
organization_id = int(organization.id)
properties = {
"organization_id": organization_id,
"organization_provider_id": getattr(organization, "provider_id", None),
"auth_provider": "stack",
}
if uses_mps_billing_v2 is not None:
properties[POSTHOG_ORGANIZATION_USES_MPS_BILLING_V2_PROPERTY] = (
uses_mps_billing_v2
)
group_identify(
POSTHOG_ORGANIZATION_GROUP_TYPE,
str(organization_id),
properties,
)
except Exception:
logger.exception("Failed to sync organization group properties to PostHog")
def _sync_posthog_organization_mps_billing_v2_status(
organization_id: int,
*,
uses_mps_billing_v2: bool,
) -> None:
"""Update the PostHog organization group with current MPS billing status."""
try:
organization_id = int(organization_id)
group_identify(
POSTHOG_ORGANIZATION_GROUP_TYPE,
str(organization_id),
{POSTHOG_ORGANIZATION_USES_MPS_BILLING_V2_PROPERTY: uses_mps_billing_v2},
)
except Exception:
logger.exception("Failed to sync organization billing status to PostHog")
def _associate_user_with_posthog_organization(
*,
user: UserModel,

View file

@ -64,6 +64,7 @@ class UserConfigurationValidator:
ServiceProviders.RIME.value: self._check_rime_api_key,
ServiceProviders.MINIMAX.value: self._check_minimax_api_key,
ServiceProviders.SMALLEST.value: self._check_smallest_api_key,
ServiceProviders.XAI.value: self._check_xai_api_key,
}
async def validate(
@ -376,6 +377,32 @@ class UserConfigurationValidator:
def _check_grok_realtime_api_key(self, model: str, api_key: str) -> bool:
return True
def _check_xai_api_key(self, model: str, api_key: str) -> bool:
# Use the TTS voices endpoint as a best-effort smoke test. Some xAI keys
# can be scoped in ways that block listing voices even though the key is
# still intended for TTS usage, so only a clear auth failure rejects save.
try:
response = httpx.get(
"https://api.x.ai/v1/tts/voices",
headers={"Authorization": f"Bearer {api_key}"},
timeout=10.0,
)
except httpx.RequestError:
raise ValueError(
"Could not connect to the xAI API. Please check your network "
"connection and try again."
)
if response.status_code == 200:
return True
if response.status_code == 401:
raise ValueError(
"Invalid xAI API key. The key was rejected by the xAI API. "
"Please check that your API key is correct and active. "
"You can verify your keys at "
"https://console.x.ai."
)
return True
def _check_ultravox_realtime_api_key(self, model: str, api_key: str) -> bool:
return True

View file

@ -92,6 +92,7 @@ class ServiceProviders(str, Enum):
GOOGLE_VERTEX_REALTIME = "google_vertex_realtime"
AZURE_REALTIME = "azure_realtime"
SMALLEST = "smallest"
XAI = "xai"
class BaseServiceConfiguration(BaseModel):
@ -122,6 +123,7 @@ class BaseServiceConfiguration(BaseModel):
ServiceProviders.AZURE_REALTIME,
ServiceProviders.SARVAM,
ServiceProviders.SMALLEST,
ServiceProviders.XAI,
]
api_key: str | list[str]
@ -256,6 +258,7 @@ GOOGLE_VERTEX_REALTIME_PROVIDER_MODEL_CONFIG = provider_model_config(
DEEPGRAM_PROVIDER_MODEL_CONFIG = provider_model_config("Deepgram")
ELEVENLABS_PROVIDER_MODEL_CONFIG = provider_model_config("ElevenLabs")
CARTESIA_PROVIDER_MODEL_CONFIG = provider_model_config("Cartesia")
XAI_PROVIDER_MODEL_CONFIG = provider_model_config("xAI")
INWORLD_PROVIDER_MODEL_CONFIG = provider_model_config(
"Inworld",
description=(
@ -531,6 +534,7 @@ class HuggingFaceLLMConfiguration(BaseLLMConfiguration):
MINIMAX_MODELS = [
"MiniMax-M2.7",
"MiniMax-M2.7-highspeed",
"MiniMax-M3",
]
@ -1278,6 +1282,32 @@ class SmallestAITTSConfiguration(BaseTTSConfiguration):
)
XAI_TTS_VOICES = ["eve", "ara", "leo", "rex", "sal"]
@register_tts
class XAITTSConfiguration(BaseServiceConfiguration):
model_config = XAI_PROVIDER_MODEL_CONFIG
provider: Literal[ServiceProviders.XAI] = ServiceProviders.XAI
voice: str = Field(
default="eve",
description="xAI voice persona.",
json_schema_extra={"examples": XAI_TTS_VOICES, "allow_custom_input": True},
)
language: str = Field(
default="en",
description="BCP-47 language code for synthesis (e.g. 'en', 'fr', 'de'), or 'auto' for automatic language detection.",
json_schema_extra={"allow_custom_input": True},
)
@computed_field
@property
def model(self) -> str:
# xAI TTS has no separate model selector; the voice fully specifies the
# output. A constant keeps the shared `.model` contract satisfied.
return "xai-tts"
TTSConfig = Annotated[
Union[
DeepgramTTSConfiguration,
@ -1294,6 +1324,7 @@ TTSConfig = Annotated[
MiniMaxTTSConfiguration,
AzureSpeechTTSConfiguration,
SmallestAITTSConfiguration,
XAITTSConfiguration,
],
Field(discriminator="provider"),
]

View file

@ -2,9 +2,8 @@
Centralizes the provider branching (Azure BYOK / Dograh-managed / OpenAI-compatible
BYOK) that was previously duplicated across document ingestion, the search route,
and the RAG tool, and resolves the MPS billing v2 protocol the same way the voice
path does: attach it only for orgs already on v2, and never create a billing
account to do so.
and the RAG tool, and resolves the MPS correlation id the same way the voice
path does.
"""
from typing import Optional
@ -24,46 +23,23 @@ DEFAULT_AZURE_API_VERSION = "2024-02-15-preview"
async def resolve_embedding_correlation_id(
*,
organization_id: Optional[int],
service_key: Optional[str],
created_by: Optional[str] = None,
) -> Optional[str]:
"""Resolve an MPS correlation id for a managed embedding call made outside a run.
"""Mint an MPS correlation id for a managed embedding call made outside a run.
Mirrors the voice path's gating:
- OSS deployments use a pasted hosted v2 key (v2 by definition), so mint
directly via the bearer endpoint matching ``_authorize_oss_managed_v2_correlation``.
- Hosted/SaaS: read the org's billing mode (no side effects) and mint only when
it is already v2. Minting for an already-v2 org is a no-op on the account.
Returns ``None`` when the call should be sent without the protocol; MPS accepts
un-gated embedding calls from v1 orgs. Never creates a v2 billing account.
Matches the voice path's ``_authorize_oss_managed_v2_correlation``: the
correlation is minted via the bearer service-key endpoint, so it works for
hosted orgs and OSS keys alike. Returns ``None`` when minting fails; MPS
accepts un-correlated embedding calls.
"""
if not service_key:
return None
# Imported lazily to avoid import-time cycles between the gen_ai and service
# layers (matches the inline-import convention used elsewhere in the app).
from api.constants import DEPLOYMENT_MODE
from api.services.mps_service_key_client import mps_service_key_client
try:
if DEPLOYMENT_MODE == "oss":
minted = await mps_service_key_client.create_correlation_id(
service_key=service_key
)
return minted.get("correlation_id")
if organization_id is None:
return None
status = await mps_service_key_client.get_billing_account_status(
organization_id, created_by=created_by
)
if not status or status.get("billing_mode") != "v2":
return None
minted = await mps_service_key_client.create_correlation_id(
service_key=service_key
)
@ -71,7 +47,7 @@ async def resolve_embedding_correlation_id(
except Exception as e:
logger.warning(
"Could not resolve MPS correlation id for managed embeddings; "
"sending without v2 protocol: {}",
"sending without it: {}",
e,
)
return None
@ -87,8 +63,6 @@ async def build_embedding_service(
endpoint: Optional[str] = None,
api_version: Optional[str] = None,
correlation_id: Optional[str] = None,
organization_id: Optional[int] = None,
created_by: Optional[str] = None,
resolve_correlation: bool = False,
) -> BaseEmbeddingService:
"""Construct the right embedding service for a provider/config.
@ -116,11 +90,7 @@ async def build_embedding_service(
if provider == ServiceProviders.DOGRAH.value:
cid = correlation_id
if cid is None and resolve_correlation:
cid = await resolve_embedding_correlation_id(
organization_id=organization_id,
service_key=api_key,
created_by=created_by,
)
cid = await resolve_embedding_correlation_id(service_key=api_key)
return DograhEmbeddingService(
db_client=db_client,
api_key=api_key,

View file

@ -241,19 +241,12 @@ class MPSServiceKeyClient:
)
return False
async def check_service_key_usage(
self,
service_key: str,
organization_id: Optional[int] = None,
created_by: Optional[str] = None,
) -> dict:
async def check_service_key_usage(self, service_key: str) -> dict:
"""
Check the usage and quota of a service key.
Args:
service_key: The service key to check usage for
organization_id: Organization ID (for authenticated mode)
created_by: User provider ID (for OSS mode)
Returns:
Dictionary containing:
@ -321,39 +314,6 @@ class MPSServiceKeyClient:
response=response,
)
async def get_usage_by_organization(self, organization_id: int) -> dict:
"""
Get aggregated usage for all service keys belonging to an organization (hosted mode).
Args:
organization_id: The organization's ID
Returns:
Dictionary containing total_credits_used and remaining_credits
"""
async with httpx.AsyncClient(timeout=self.timeout) as client:
response = await client.post(
f"{self.base_url}/api/v1/service-keys/usage/organization",
json={"organization_id": organization_id},
headers=self._get_headers(organization_id=organization_id),
)
if response.status_code == 200:
data = response.json()
return {
"total_credits_used": data.get("total_credits_used", 0.0),
"remaining_credits": data.get("remaining_credits", 0.0),
}
else:
logger.error(
f"Failed to get usage by organization: {response.status_code} - {response.text}"
)
raise httpx.HTTPStatusError(
f"Failed to get usage by organization: {response.text}",
request=response.request,
response=response,
)
async def create_credit_purchase_url(
self,
organization_id: int,
@ -422,34 +382,6 @@ class MPSServiceKeyClient:
response=response,
)
async def get_billing_account_status(
self,
organization_id: int,
created_by: Optional[str] = None,
) -> Optional[dict]:
"""Get an existing MPS v2 billing account without creating one."""
async with httpx.AsyncClient(timeout=self.timeout) as client:
response = await client.get(
f"{self.base_url}/api/v1/billing/accounts/{organization_id}/status",
headers=self._get_headers(
organization_id=organization_id,
created_by=created_by,
),
)
if response.status_code == 200:
return response.json()
logger.error(
"Failed to get MPS billing account status: "
f"{response.status_code} - {response.text}"
)
raise httpx.HTTPStatusError(
f"Failed to get MPS billing account status: {response.text}",
request=response.request,
response=response,
)
async def ensure_billing_account_v2(
self,
organization_id: int,

View file

@ -70,6 +70,7 @@ def register_event_handlers(
pre_call_fetch_task: asyncio.Task | None = None,
user_provider_id: str | None = None,
integration_runtime_sessions: list[IntegrationRuntimeSession] | None = None,
include_transcript_end_timestamps: bool = False,
):
"""Register all event handlers for transport and task events.
@ -386,7 +387,9 @@ def register_event_handlers(
else:
logger.debug("Bot audio buffer is empty, skipping upload")
transcript_text = in_memory_logs_buffer.generate_transcript_text()
transcript_text = in_memory_logs_buffer.generate_transcript_text(
include_end_timestamps=include_transcript_end_timestamps
)
if not transcript_text:
logger.debug("No transcript events in logs buffer, skipping upload")

View file

@ -107,6 +107,12 @@ class InMemoryLogsBuffer:
self._turn_counter = 0
self._current_node_id: Optional[str] = None
self._current_node_name: Optional[str] = None
self._user_speech_start_timestamp: Optional[str] = None
self._user_speech_end_timestamp: Optional[str] = None
self._user_speech_start_from_vad = False
self._user_speech_end_from_vad = False
self._bot_speech_start_timestamp: Optional[str] = None
self._bot_speech_end_timestamp: Optional[str] = None
def set_current_node(self, node_id: str, node_name: str):
"""Set the current node ID and name to be injected into subsequent events."""
@ -123,11 +129,126 @@ class InMemoryLogsBuffer:
"""Get the current node name."""
return self._current_node_name
@staticmethod
def _now_iso() -> str:
return datetime.now(UTC).isoformat(timespec="milliseconds")
def mark_user_started_speaking(
self, timestamp: Optional[str] = None, *, from_vad: bool = False
):
"""Record when the user started speaking for the current turn."""
vad_interval_is_open = (
self._user_speech_start_from_vad and self._user_speech_end_timestamp is None
)
if vad_interval_is_open and not from_vad:
return
self._user_speech_start_timestamp = timestamp or self._now_iso()
self._user_speech_end_timestamp = None
self._user_speech_start_from_vad = from_vad
self._user_speech_end_from_vad = False
self._update_latest_payload_start_timestamp(
RealtimeFeedbackType.USER_TRANSCRIPTION.value,
self._user_speech_start_timestamp,
require_final=True,
)
def mark_user_stopped_speaking(
self, timestamp: Optional[str] = None, *, from_vad: bool = False
):
"""Record when the user stopped speaking and update the latest user event."""
if self._user_speech_end_from_vad and not from_vad:
return
self._user_speech_end_timestamp = timestamp or self._now_iso()
self._user_speech_end_from_vad = from_vad
self._update_latest_payload_end_timestamp(
RealtimeFeedbackType.USER_TRANSCRIPTION.value,
self._user_speech_end_timestamp,
require_final=True,
)
def mark_bot_started_speaking(self, timestamp: Optional[str] = None):
"""Record when the bot started speaking for the current assistant turn."""
self._bot_speech_start_timestamp = timestamp or self._now_iso()
self._bot_speech_end_timestamp = None
self._update_latest_payload_start_timestamp(
RealtimeFeedbackType.BOT_TEXT.value,
self._bot_speech_start_timestamp,
)
def mark_bot_stopped_speaking(self, timestamp: Optional[str] = None):
"""Record when the bot stopped speaking and update the latest bot event."""
self._bot_speech_end_timestamp = timestamp or self._now_iso()
self._update_latest_payload_end_timestamp(
RealtimeFeedbackType.BOT_TEXT.value,
self._bot_speech_end_timestamp,
)
def _find_latest_open_payload(
self, event_type: str, *, require_final: bool = False
) -> dict | None:
for event in reversed(self._events):
if event.get("type") != event_type:
continue
payload = event.get("payload")
if not isinstance(payload, dict):
continue
if require_final and payload.get("final") is not True:
continue
if payload.get("end_timestamp"):
continue
return payload
return None
def _update_latest_payload_start_timestamp(
self, event_type: str, start_timestamp: str, *, require_final: bool = False
):
payload = self._find_latest_open_payload(
event_type, require_final=require_final
)
if payload is not None:
payload["timestamp"] = start_timestamp
def _update_latest_payload_end_timestamp(
self, event_type: str, end_timestamp: str, *, require_final: bool = False
):
payload = self._find_latest_open_payload(
event_type, require_final=require_final
)
if payload is not None:
payload["end_timestamp"] = end_timestamp
def _event_with_speech_timestamps(self, event: dict) -> dict:
event_type = event.get("type")
payload = event.get("payload")
if not isinstance(payload, dict):
return event
payload_with_timestamps = dict(payload)
if (
event_type == RealtimeFeedbackType.USER_TRANSCRIPTION.value
and payload.get("final") is True
):
if self._user_speech_start_timestamp:
payload_with_timestamps["timestamp"] = self._user_speech_start_timestamp
if self._user_speech_end_timestamp:
payload_with_timestamps["end_timestamp"] = self._user_speech_end_timestamp
elif event_type == RealtimeFeedbackType.BOT_TEXT.value:
bot_interval_is_active = self._bot_speech_end_timestamp is None
if bot_interval_is_active and self._bot_speech_start_timestamp:
payload_with_timestamps["timestamp"] = self._bot_speech_start_timestamp
if payload_with_timestamps == payload:
return event
return {**event, "payload": payload_with_timestamps}
async def append(self, event: dict):
"""Append a feedback event to the buffer with timestamp and current node."""
event = self._event_with_speech_timestamps(event)
timestamped_event = stamp_realtime_feedback_event(
event,
timestamp=datetime.now(UTC).isoformat(),
timestamp=self._now_iso(),
turn=self._turn_counter,
node_id=self._current_node_id,
node_name=self._current_node_name,
@ -166,13 +287,15 @@ class InMemoryLogsBuffer:
return True
return False
def generate_transcript_text(self) -> str:
def generate_transcript_text(self, *, include_end_timestamps: bool = False) -> str:
"""Generate transcript text from logged events.
Filters for rtf-user-transcription (final) and rtf-bot-text events,
formats them as '[timestamp] user/assistant: text\\n'.
"""
return _generate_transcript_text(self._sorted_events())
return _generate_transcript_text(
self._sorted_events(), include_end_timestamps=include_end_timestamps
)
@property
def is_empty(self) -> bool:

View file

@ -30,6 +30,7 @@ def build_user_transcription_event(
text: str,
final: bool,
timestamp: str | None = None,
end_timestamp: str | None = None,
user_id: str | None = None,
) -> dict[str, Any]:
payload: dict[str, Any] = {
@ -38,6 +39,8 @@ def build_user_transcription_event(
}
if timestamp is not None:
payload["timestamp"] = timestamp
if end_timestamp is not None:
payload["end_timestamp"] = end_timestamp
if user_id is not None:
payload["user_id"] = user_id
return {
@ -50,10 +53,13 @@ def build_bot_text_event(
*,
text: str,
timestamp: str | None = None,
end_timestamp: str | None = None,
) -> dict[str, Any]:
payload: dict[str, Any] = {"text": text}
if timestamp is not None:
payload["timestamp"] = timestamp
if end_timestamp is not None:
payload["end_timestamp"] = end_timestamp
return {
"type": RealtimeFeedbackType.BOT_TEXT.value,
"payload": payload,

View file

@ -21,6 +21,7 @@ node changes.
"""
import json
from datetime import UTC, datetime
from typing import TYPE_CHECKING, Awaitable, Callable, Optional, Set
from loguru import logger
@ -52,8 +53,12 @@ from pipecat.frames.frames import (
TranscriptionFrame,
TTSSpeakFrame,
TTSTextFrame,
UserStartedSpeakingFrame,
UserStoppedSpeakingFrame,
UserMuteStartedFrame,
UserMuteStoppedFrame,
VADUserStartedSpeakingFrame,
VADUserStoppedSpeakingFrame,
)
from pipecat.metrics.metrics import TTFBMetricsData
from pipecat.observers.base_observer import BaseObserver, FramePushed
@ -62,6 +67,10 @@ from pipecat.transports.base_output import BaseOutputTransport
from pipecat.utils.enums import RealtimeFeedbackType
def _epoch_seconds_to_utc_iso(timestamp: float) -> str:
return datetime.fromtimestamp(timestamp, UTC).isoformat(timespec="milliseconds")
class RealtimeFeedbackObserver(BaseObserver):
"""Observer that sends real-time events via WebSocket and persists final transcripts.
@ -138,13 +147,35 @@ class RealtimeFeedbackObserver(BaseObserver):
return
# Bot speaking state - WS only (ephemeral state signals, not persisted)
elif isinstance(frame, BotStartedSpeakingFrame):
if self._logs_buffer:
self._logs_buffer.mark_bot_started_speaking()
await self._send_ws(
{"type": RealtimeFeedbackType.BOT_STARTED_SPEAKING.value, "payload": {}}
)
elif isinstance(frame, BotStoppedSpeakingFrame):
if self._logs_buffer:
self._logs_buffer.mark_bot_stopped_speaking()
await self._send_ws(
{"type": RealtimeFeedbackType.BOT_STOPPED_SPEAKING.value, "payload": {}}
)
elif isinstance(frame, UserStartedSpeakingFrame):
if self._logs_buffer:
self._logs_buffer.mark_user_started_speaking()
elif isinstance(frame, UserStoppedSpeakingFrame):
if self._logs_buffer:
self._logs_buffer.mark_user_stopped_speaking()
elif isinstance(frame, VADUserStartedSpeakingFrame):
if self._logs_buffer:
self._logs_buffer.mark_user_started_speaking(
_epoch_seconds_to_utc_iso(frame.timestamp - frame.start_secs),
from_vad=True,
)
elif isinstance(frame, VADUserStoppedSpeakingFrame):
if self._logs_buffer:
self._logs_buffer.mark_user_stopped_speaking(
_epoch_seconds_to_utc_iso(frame.timestamp - frame.stop_secs),
from_vad=True,
)
# User mute state - WS only (ephemeral state signals, not persisted)
elif isinstance(frame, UserMuteStartedFrame):
await self._send_ws(
@ -314,6 +345,7 @@ def register_turn_log_handlers(
text=message.content,
final=True,
timestamp=message.timestamp,
end_timestamp=getattr(message, "end_timestamp", None),
)
)
except Exception as e:
@ -327,6 +359,7 @@ def register_turn_log_handlers(
build_bot_text_event(
text=message.content,
timestamp=message.timestamp,
end_timestamp=getattr(message, "end_timestamp", None),
)
)
except Exception as e:

View file

@ -557,6 +557,10 @@ async def _run_pipeline_impl(
max_call_duration_seconds = DEFAULT_MAX_CALL_DURATION_SECONDS
max_user_idle_timeout = DEFAULT_MAX_USER_IDLE_TIMEOUT_SECONDS
keyterms = None # Dictionary words for STT boosting
transcript_config = run_configs.get("transcript_configuration") or {}
include_transcript_end_timestamps = bool(
transcript_config.get("include_end_timestamps", False)
)
if run_configs:
if "max_call_duration" in run_configs:
@ -1057,6 +1061,7 @@ async def _run_pipeline_impl(
pre_call_fetch_task=pre_call_fetch_task,
user_provider_id=user_provider_id,
integration_runtime_sessions=integration_runtime_sessions,
include_transcript_end_timestamps=include_transcript_end_timestamps,
)
register_audio_data_handler(audio_buffer, workflow_run_id, in_memory_audio_buffer)

View file

@ -81,6 +81,7 @@ from pipecat.services.speechmatics.stt import (
SpeechmaticsSTTService,
SpeechmaticsSTTSettings,
)
from pipecat.services.xai.tts import XAIHttpTTSService, XAITTSSettings
from pipecat.transcriptions.language import Language
from pipecat.utils.text.xml_function_tag_filter import XMLFunctionTagFilter
@ -740,6 +741,28 @@ def create_tts_service(
skip_aggregator_types=["recording_router", "recording"],
silence_time_s=1.0,
)
elif user_config.tts.provider == ServiceProviders.XAI.value:
voice = getattr(user_config.tts, "voice", None) or "eve"
language_code = getattr(user_config.tts, "language", None) or "en"
if language_code.lower() == "auto":
pipecat_language = "auto"
else:
try:
pipecat_language = Language(language_code)
except ValueError:
pipecat_language = Language.EN
return XAIHttpTTSService(
api_key=user_config.tts.api_key,
sample_rate=audio_config.transport_out_sample_rate,
encoding="pcm",
settings=XAITTSSettings(
voice=voice,
language=pipecat_language,
),
text_filters=[xml_function_tag_filter],
skip_aggregator_types=["recording_router", "recording"],
silence_time_s=1.0,
)
else:
raise HTTPException(
status_code=400, detail=f"Invalid TTS provider {user_config.tts.provider}"

View file

@ -25,13 +25,13 @@ from api.services.mps_service_key_client import mps_service_key_client
MINIMUM_DOGRAH_CREDITS_FOR_CALL = 0.10
LEGACY_QUOTA_EXCEEDED_MESSAGE = (
OSS_QUOTA_EXCEEDED_MESSAGE = (
"You have exhausted your trial credits. "
"Please email founders@dograh.com for additional Dograh credits "
"or change providers in Models configurations."
"Please sign up on app.dograh.com to create a "
"new service key and set up in your model configurations."
)
BILLING_V2_QUOTA_EXCEEDED_MESSAGE = (
HOSTED_QUOTA_EXCEEDED_MESSAGE = (
"You have exhausted your Dograh credits. "
"Please purchase more credits from /billing "
"or change providers in Models configurations."
@ -60,19 +60,19 @@ def _safe_float(value: Any, default: float = 0.0) -> float:
return default
def _insufficient_billing_v2_quota_result() -> QuotaCheckResult:
def _insufficient_hosted_quota_result() -> QuotaCheckResult:
return QuotaCheckResult(
has_quota=False,
error_code="insufficient_credits",
error_message=BILLING_V2_QUOTA_EXCEEDED_MESSAGE,
error_message=HOSTED_QUOTA_EXCEEDED_MESSAGE,
)
def _insufficient_legacy_quota_result() -> QuotaCheckResult:
def _insufficient_oss_quota_result() -> QuotaCheckResult:
return QuotaCheckResult(
has_quota=False,
error_code="quota_exceeded",
error_message=LEGACY_QUOTA_EXCEEDED_MESSAGE,
error_message=OSS_QUOTA_EXCEEDED_MESSAGE,
)
@ -157,10 +157,10 @@ async def _authorize_hosted_workflow_run_start(
workflow_id: int | None,
workflow_run_id: int | None,
user_config: Any,
) -> tuple[QuotaCheckResult, bool]:
"""Authorize hosted v2 billing and return whether MPS handled enforcement."""
if DEPLOYMENT_MODE == "oss" or organization_id is None:
return QuotaCheckResult(has_quota=True), False
) -> QuotaCheckResult:
"""Authorize a hosted workflow run against the org's MPS billing account."""
if organization_id is None:
return QuotaCheckResult(has_quota=True)
requires_correlation = bool(
workflow_run_id and uses_managed_model_services_v2(user_config)
@ -169,16 +169,13 @@ async def _authorize_hosted_workflow_run_start(
get_dograh_service_api_key(user_config) if requires_correlation else None
)
if requires_correlation and not service_key:
return (
QuotaCheckResult(
has_quota=False,
error_code="invalid_service_key",
error_message=(
"You have invalid keys in your model configuration. "
"Please validate the service keys."
),
return QuotaCheckResult(
has_quota=False,
error_code="invalid_service_key",
error_message=(
"You have invalid keys in your model configuration. "
"Please validate the service keys."
),
True,
)
try:
@ -205,38 +202,28 @@ async def _authorize_hosted_workflow_run_start(
e,
)
if _is_service_key_org_mismatch_error(e):
return (
QuotaCheckResult(
has_quota=False,
error_code="service_key_org_mismatch",
error_message=SERVICE_TOKEN_ORG_MISMATCH_MESSAGE,
),
True,
)
return (
QuotaCheckResult(
return QuotaCheckResult(
has_quota=False,
error_code="quota_check_failed",
error_message="Could not verify Dograh credits. Please try again.",
),
True,
error_code="service_key_org_mismatch",
error_message=SERVICE_TOKEN_ORG_MISMATCH_MESSAGE,
)
return QuotaCheckResult(
has_quota=False,
error_code="quota_check_failed",
error_message="Could not verify Dograh credits. Please try again.",
)
billing_mode = authorization.get("billing_mode")
if billing_mode != "v2":
return QuotaCheckResult(has_quota=True), False
remaining = _safe_float(authorization.get("remaining_credits"))
if (
not authorization.get("allowed", False)
or remaining < MINIMUM_DOGRAH_CREDITS_FOR_CALL
):
logger.warning(
"Insufficient Dograh billing v2 credits for org {}: {:.2f} credits remaining",
"Insufficient Dograh credits for org {}: {:.2f} credits remaining",
organization_id,
remaining,
)
return _insufficient_billing_v2_quota_result(), True
return _insufficient_hosted_quota_result()
try:
await _store_run_correlation_id(
@ -249,35 +236,27 @@ async def _authorize_hosted_workflow_run_start(
workflow_run_id,
e,
)
return (
QuotaCheckResult(
has_quota=False,
error_code="quota_check_failed",
error_message="Could not verify Dograh credits. Please try again.",
),
True,
return QuotaCheckResult(
has_quota=False,
error_code="quota_check_failed",
error_message="Could not verify Dograh credits. Please try again.",
)
logger.info(
"Dograh billing v2 run authorization passed for org {}: {:.2f} credits remaining",
"Dograh run authorization passed for org {}: {:.2f} credits remaining",
organization_id,
remaining,
)
return QuotaCheckResult(has_quota=True), True
return QuotaCheckResult(has_quota=True)
async def _authorize_legacy_dograh_keys(
async def _authorize_oss_dograh_keys(
*,
dograh_api_keys: set[str],
organization_id: int | None,
workflow_owner: UserModel,
) -> QuotaCheckResult:
"""Check per-key MPS credits for OSS deployments before a run starts."""
for api_key in dograh_api_keys:
try:
usage = await mps_service_key_client.check_service_key_usage(
api_key,
organization_id=organization_id,
created_by=workflow_owner.provider_id,
)
usage = await mps_service_key_client.check_service_key_usage(api_key)
remaining = usage.get("remaining_credits", 0.0)
# Require at least $0.10 for a short call
@ -286,7 +265,7 @@ async def _authorize_legacy_dograh_keys(
f"Insufficient Dograh credits for key ...{api_key[-8:]}: "
f"${remaining:.2f} remaining"
)
return _insufficient_legacy_quota_result()
return _insufficient_oss_quota_result()
logger.info(
f"Dograh quota check passed for key ...{api_key[-8:]}: "
@ -363,9 +342,9 @@ async def authorize_workflow_run_start(
) -> QuotaCheckResult:
"""Authorize a workflow run before any billable call/text runtime starts.
The workflow organization is the billing subject for hosted v2. The workflow
owner is used only to resolve the effective model configuration and legacy
service-key metadata.
The workflow organization is the billing subject for hosted deployments.
OSS deployments are billed per service key instead. The workflow owner is
used only to resolve the effective model configuration.
"""
try:
workflow = await db_client.get_workflow_by_id(workflow_id)
@ -376,19 +355,38 @@ async def authorize_workflow_run_start(
error_message="Workflow not found",
)
actor_org_id = getattr(actor_user, "selected_organization_id", None)
if actor_org_id is not None and actor_org_id != workflow.organization_id:
logger.warning(
"Workflow start authorization denied: actor org {} does not match workflow {} org {}",
actor_org_id,
workflow_id,
workflow.organization_id,
)
return QuotaCheckResult(
has_quota=False,
error_code="workflow_not_found",
error_message="Workflow not found",
)
actor_id = getattr(actor_user, "id", None)
if actor_id is not None and workflow.organization_id is not None:
try:
is_member = await db_client.is_user_member_of_organization(
user_id=actor_id,
organization_id=workflow.organization_id,
)
except Exception as e:
logger.error(
"Workflow start authorization denied: failed to validate actor {} membership for workflow {} org {}: {}",
actor_id,
workflow_id,
workflow.organization_id,
e,
)
return QuotaCheckResult(
has_quota=False,
error_code="workflow_not_found",
error_message="Workflow not found",
)
if not is_member:
logger.warning(
"Workflow start authorization denied: actor {} is not a member of workflow {} org {}",
actor_id,
workflow_id,
workflow.organization_id,
)
return QuotaCheckResult(
has_quota=False,
error_code="workflow_not_found",
error_message="Workflow not found",
)
workflow_owner = await db_client.get_user_by_id(workflow.user_id)
if not workflow_owner:
@ -405,38 +403,27 @@ async def authorize_workflow_run_start(
)
if DEPLOYMENT_MODE != "oss":
hosted_result, hosted_enforced = await _authorize_hosted_workflow_run_start(
return await _authorize_hosted_workflow_run_start(
workflow_owner=workflow_owner,
organization_id=workflow.organization_id,
workflow_id=workflow.id,
workflow_run_id=workflow_run_id,
user_config=user_config,
)
if hosted_enforced or not hosted_result.has_quota:
return hosted_result
dograh_api_keys = _dograh_api_keys(user_config)
if not dograh_api_keys:
return QuotaCheckResult(has_quota=True)
legacy_result = await _authorize_legacy_dograh_keys(
dograh_api_keys=dograh_api_keys,
organization_id=(
None if DEPLOYMENT_MODE == "oss" else workflow.organization_id
),
workflow_owner=workflow_owner,
)
if not legacy_result.has_quota:
return legacy_result
if DEPLOYMENT_MODE == "oss":
return await _authorize_oss_managed_v2_correlation(
workflow_id=workflow.id,
workflow_run_id=workflow_run_id,
user_config=user_config,
if dograh_api_keys:
oss_result = await _authorize_oss_dograh_keys(
dograh_api_keys=dograh_api_keys,
)
if not oss_result.has_quota:
return oss_result
return QuotaCheckResult(has_quota=True)
return await _authorize_oss_managed_v2_correlation(
workflow_id=workflow.id,
workflow_run_id=workflow_run_id,
user_config=user_config,
)
except Exception as e:
logger.error(f"Error during quota check: {str(e)}")

View file

@ -1,10 +1,8 @@
"""Rule-based audit of a workflow definition's nodes + edges.
Pure, dependency-free helpers derived from `NodeSpec.graph_constraints`.
Lives in tracked code so the regression tests in
`test_workflow_graph_constraints.py` can pin it; the admin cleanup
script in `api/services/admin_utils/local_exec.py` is the production
consumer.
Lives in tracked code so `test_workflow_graph_constraints.py` can pin the
verdicts that one-off cleanup tooling needs to share with runtime validation.
"""
from collections import Counter

View file

@ -266,8 +266,7 @@ async def _perform_retrieval(
)
# Search runs inside a workflow run: reuse the run's MPS correlation
# id (present only for v2 orgs; None otherwise → sent without the
# protocol). The Dograh-managed path forwards it via request metadata.
# id. The Dograh-managed path forwards it via request metadata.
embedding_service = await build_embedding_service(
db_client=db_client,
provider=embeddings_provider,

View file

@ -32,13 +32,6 @@ def _duration_seconds_from_usage_info(workflow_run) -> float | None:
return duration_seconds if duration_seconds > 0 else None
async def _organization_uses_mps_billing_v2(organization_id: int) -> bool:
account = await mps_service_key_client.get_billing_account_status(
organization_id=organization_id
)
return bool(account and account.get("billing_mode") == "v2")
def _is_usage_not_ready_error(exc: Exception) -> bool:
response = getattr(exc, "response", None)
if getattr(response, "status_code", None) != 409:
@ -79,12 +72,6 @@ async def report_workflow_run_platform_usage(workflow_run) -> None:
return
try:
if not await _organization_uses_mps_billing_v2(organization_id):
logger.debug(
"Not reporting platform usage since org not using mps billing v2"
)
return
result = await mps_service_key_client.report_platform_usage(
organization_id=organization_id,
correlation_id=correlation_id,