mirror of
https://github.com/dograh-hq/dograh.git
synced 2026-07-25 12:01:04 +02:00
Merge remote-tracking branch 'origin/main' into fix/org-scoped-access
# Conflicts: # api/routes/agent_stream.py # api/routes/telephony.py # api/routes/webrtc_signaling.py # docs/api-reference/openapi.json # sdk/python/src/dograh_sdk/_generated_models.py
This commit is contained in:
commit
d22c073cb5
38 changed files with 2332 additions and 471 deletions
|
|
@ -5,10 +5,14 @@ from typing import TYPE_CHECKING, Optional
|
|||
|
||||
from loguru import logger
|
||||
|
||||
from api.constants import DEFAULT_ORG_CONCURRENCY_LIMIT
|
||||
from api.db import db_client
|
||||
from api.db.models import QueuedRunModel, WorkflowRunModel
|
||||
from api.enums import OrganizationConfigurationKey, WorkflowRunState
|
||||
from api.enums import WorkflowRunState
|
||||
from api.services.call_concurrency import (
|
||||
CallConcurrencyLimitError,
|
||||
CallConcurrencySlot,
|
||||
call_concurrency,
|
||||
)
|
||||
from api.services.campaign.circuit_breaker import circuit_breaker
|
||||
from api.services.campaign.errors import (
|
||||
ConcurrentSlotAcquisitionError,
|
||||
|
|
@ -29,9 +33,6 @@ if TYPE_CHECKING:
|
|||
class CampaignCallDispatcher:
|
||||
"""Manages rate-limited and concurrent-limited call dispatching"""
|
||||
|
||||
def __init__(self):
|
||||
self.default_concurrent_limit = int(DEFAULT_ORG_CONCURRENCY_LIMIT)
|
||||
|
||||
async def get_provider_for_campaign(self, campaign) -> "TelephonyProvider":
|
||||
"""Get the telephony provider pinned to this campaign's config. Falls back
|
||||
to the org's default config for legacy campaigns whose
|
||||
|
|
@ -53,18 +54,7 @@ class CampaignCallDispatcher:
|
|||
|
||||
async def get_org_concurrent_limit(self, organization_id: int) -> int:
|
||||
"""Get the concurrent call limit for an organization."""
|
||||
try:
|
||||
config = await db_client.get_configuration(
|
||||
organization_id,
|
||||
OrganizationConfigurationKey.CONCURRENT_CALL_LIMIT.value,
|
||||
)
|
||||
if config and config.value:
|
||||
return int(config.value["value"])
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Error getting concurrent limit for org {organization_id}: {e}"
|
||||
)
|
||||
return self.default_concurrent_limit
|
||||
return await call_concurrency.get_org_concurrent_limit(organization_id)
|
||||
|
||||
async def process_batch(self, campaign_id: int, batch_size: int = 10) -> int:
|
||||
"""
|
||||
|
|
@ -119,12 +109,14 @@ class CampaignCallDispatcher:
|
|||
)
|
||||
|
||||
# Acquire concurrent slot - waits until a slot is available
|
||||
slot_id = await self.acquire_concurrent_slot(
|
||||
concurrency_slot = await self.acquire_concurrent_slot(
|
||||
campaign.organization_id, campaign
|
||||
)
|
||||
|
||||
# Dispatch the call
|
||||
workflow_run = await self.dispatch_call(queued_run, campaign, slot_id)
|
||||
workflow_run = await self.dispatch_call(
|
||||
queued_run, campaign, concurrency_slot
|
||||
)
|
||||
|
||||
# Update queued run as processed
|
||||
await db_client.update_queued_run(
|
||||
|
|
@ -233,68 +225,61 @@ class CampaignCallDispatcher:
|
|||
)
|
||||
|
||||
async def dispatch_call(
|
||||
self, queued_run: QueuedRunModel, campaign: any, slot_id: str
|
||||
self,
|
||||
queued_run: QueuedRunModel,
|
||||
campaign: any,
|
||||
concurrency_slot: CallConcurrencySlot,
|
||||
) -> Optional[WorkflowRunModel]:
|
||||
"""Creates workflow run and initiates call. Requires a pre-acquired slot_id."""
|
||||
"""Creates workflow run and initiates call. Requires a pre-acquired slot."""
|
||||
from_number = None
|
||||
workflow_run = None
|
||||
slot_bound = False
|
||||
|
||||
# Get workflow details
|
||||
workflow = await db_client.get_workflow_by_id(campaign.workflow_id)
|
||||
if not workflow:
|
||||
# Release slot before raising
|
||||
await rate_limiter.release_concurrent_slot(
|
||||
campaign.organization_id, slot_id
|
||||
)
|
||||
raise ValueError(f"Workflow {campaign.workflow_id} not found")
|
||||
|
||||
# Extract phone number
|
||||
phone_number = queued_run.context_variables.get("phone_number")
|
||||
if not phone_number:
|
||||
# Release slot before raising
|
||||
await rate_limiter.release_concurrent_slot(
|
||||
campaign.organization_id, slot_id
|
||||
)
|
||||
raise ValueError(f"No phone number in queued run {queued_run.id}")
|
||||
|
||||
# Get provider for this campaign's pinned telephony config.
|
||||
provider = await self.get_provider_for_campaign(campaign)
|
||||
workflow_run_mode = provider.PROVIDER_NAME
|
||||
|
||||
# Acquire a unique from_number from the pool scoped to this campaign's
|
||||
# telephony configuration so orgs with multiple configs don't leak
|
||||
# caller IDs across configs.
|
||||
from_number = await self.acquire_from_number(
|
||||
campaign.organization_id,
|
||||
telephony_configuration_id=campaign.telephony_configuration_id,
|
||||
)
|
||||
if from_number is None:
|
||||
# Release concurrent slot before raising
|
||||
await rate_limiter.release_concurrent_slot(
|
||||
campaign.organization_id, slot_id
|
||||
)
|
||||
raise PhoneNumberPoolExhaustedError(
|
||||
organization_id=campaign.organization_id
|
||||
)
|
||||
|
||||
logger.info(f"Provider name: {provider.PROVIDER_NAME}")
|
||||
logger.info(f"Queued run context: {queued_run.context_variables}")
|
||||
|
||||
# Merge context variables (queued_run context already includes retry info if applicable)
|
||||
initial_context = {
|
||||
**queued_run.context_variables,
|
||||
"campaign_id": campaign.id,
|
||||
"provider": provider.PROVIDER_NAME,
|
||||
"source_uuid": queued_run.source_uuid,
|
||||
"caller_number": from_number,
|
||||
"called_number": phone_number,
|
||||
"telephony_configuration_id": campaign.telephony_configuration_id,
|
||||
}
|
||||
|
||||
logger.info(f"Final initial_context: {initial_context}")
|
||||
|
||||
# Create workflow run with queued_run_id tracking
|
||||
workflow_run_name = f"WR-CAMPAIGN-{campaign.id}-{queued_run.id}"
|
||||
try:
|
||||
# Get workflow details
|
||||
workflow = await db_client.get_workflow_by_id(campaign.workflow_id)
|
||||
if not workflow:
|
||||
raise ValueError(f"Workflow {campaign.workflow_id} not found")
|
||||
|
||||
# Extract phone number
|
||||
phone_number = queued_run.context_variables.get("phone_number")
|
||||
if not phone_number:
|
||||
raise ValueError(f"No phone number in queued run {queued_run.id}")
|
||||
|
||||
# Get provider for this campaign's pinned telephony config.
|
||||
provider = await self.get_provider_for_campaign(campaign)
|
||||
workflow_run_mode = provider.PROVIDER_NAME
|
||||
|
||||
# Acquire a unique from_number from the pool scoped to this campaign's
|
||||
# telephony configuration so orgs with multiple configs don't leak
|
||||
# caller IDs across configs.
|
||||
from_number = await self.acquire_from_number(
|
||||
campaign.organization_id,
|
||||
telephony_configuration_id=campaign.telephony_configuration_id,
|
||||
)
|
||||
if from_number is None:
|
||||
raise PhoneNumberPoolExhaustedError(
|
||||
organization_id=campaign.organization_id
|
||||
)
|
||||
|
||||
logger.info(f"Provider name: {provider.PROVIDER_NAME}")
|
||||
logger.info(f"Queued run context: {queued_run.context_variables}")
|
||||
|
||||
# Merge context variables (queued_run context already includes retry info if applicable)
|
||||
initial_context = {
|
||||
**queued_run.context_variables,
|
||||
"campaign_id": campaign.id,
|
||||
"provider": provider.PROVIDER_NAME,
|
||||
"source_uuid": queued_run.source_uuid,
|
||||
"caller_number": from_number,
|
||||
"called_number": phone_number,
|
||||
"telephony_configuration_id": campaign.telephony_configuration_id,
|
||||
}
|
||||
|
||||
logger.info(f"Final initial_context: {initial_context}")
|
||||
|
||||
# Create workflow run with queued_run_id tracking
|
||||
workflow_run_name = f"WR-CAMPAIGN-{campaign.id}-{queued_run.id}"
|
||||
workflow_run = await db_client.create_workflow_run(
|
||||
name=workflow_run_name,
|
||||
workflow_id=campaign.workflow_id,
|
||||
|
|
@ -305,11 +290,8 @@ class CampaignCallDispatcher:
|
|||
queued_run_id=queued_run.id, # Link to queued run for retry tracking
|
||||
organization_id=campaign.organization_id,
|
||||
)
|
||||
|
||||
# Store slot_id mapping in Redis for cleanup later
|
||||
await rate_limiter.store_workflow_slot_mapping(
|
||||
workflow_run.id, campaign.organization_id, slot_id
|
||||
)
|
||||
await call_concurrency.bind_workflow_run(concurrency_slot, workflow_run.id)
|
||||
slot_bound = True
|
||||
|
||||
# Store from_number mapping for cleanup on call completion
|
||||
await rate_limiter.store_workflow_from_number_mapping(
|
||||
|
|
@ -320,9 +302,10 @@ class CampaignCallDispatcher:
|
|||
)
|
||||
except Exception as e:
|
||||
# Release slot and from_number on error
|
||||
await rate_limiter.release_concurrent_slot(
|
||||
campaign.organization_id, slot_id
|
||||
)
|
||||
if slot_bound and workflow_run:
|
||||
await call_concurrency.release_workflow_run_slot(workflow_run.id)
|
||||
else:
|
||||
await call_concurrency.release_slot(concurrency_slot)
|
||||
if from_number:
|
||||
await rate_limiter.release_from_number(
|
||||
campaign.organization_id,
|
||||
|
|
@ -359,21 +342,7 @@ class CampaignCallDispatcher:
|
|||
gathered_context={"error": error_message},
|
||||
)
|
||||
|
||||
mapping = await rate_limiter.get_workflow_slot_mapping(workflow_run.id)
|
||||
if mapping:
|
||||
org_id, mapped_slot_id = mapping
|
||||
await rate_limiter.release_concurrent_slot(org_id, mapped_slot_id)
|
||||
await rate_limiter.delete_workflow_slot_mapping(workflow_run.id)
|
||||
|
||||
from_number_mapping = await rate_limiter.get_workflow_from_number_mapping(
|
||||
workflow_run.id
|
||||
)
|
||||
if from_number_mapping:
|
||||
fn_org_id, fn_number, fn_tcid = from_number_mapping
|
||||
await rate_limiter.release_from_number(
|
||||
fn_org_id, fn_number, telephony_configuration_id=fn_tcid
|
||||
)
|
||||
await rate_limiter.delete_workflow_from_number_mapping(workflow_run.id)
|
||||
await self.release_call_slot(workflow_run.id)
|
||||
|
||||
raise ValueError(error_message)
|
||||
|
||||
|
|
@ -448,23 +417,7 @@ class CampaignCallDispatcher:
|
|||
reason="call_initiation_failed",
|
||||
)
|
||||
|
||||
# Release concurrent slot on failure
|
||||
mapping = await rate_limiter.get_workflow_slot_mapping(workflow_run.id)
|
||||
if mapping:
|
||||
org_id, slot_id = mapping
|
||||
await rate_limiter.release_concurrent_slot(org_id, slot_id)
|
||||
await rate_limiter.delete_workflow_slot_mapping(workflow_run.id)
|
||||
|
||||
# Release from_number on failure
|
||||
from_number_mapping = await rate_limiter.get_workflow_from_number_mapping(
|
||||
workflow_run.id
|
||||
)
|
||||
if from_number_mapping:
|
||||
fn_org_id, fn_number, fn_tcid = from_number_mapping
|
||||
await rate_limiter.release_from_number(
|
||||
fn_org_id, fn_number, telephony_configuration_id=fn_tcid
|
||||
)
|
||||
await rate_limiter.delete_workflow_from_number_mapping(workflow_run.id)
|
||||
await self.release_call_slot(workflow_run.id)
|
||||
|
||||
raise
|
||||
|
||||
|
|
@ -503,7 +456,7 @@ class CampaignCallDispatcher:
|
|||
|
||||
async def acquire_concurrent_slot(
|
||||
self, organization_id: int, campaign: any, timeout: float = 600
|
||||
) -> str:
|
||||
) -> CallConcurrencySlot:
|
||||
"""
|
||||
Acquires a concurrent call slot - waits if necessary until a slot is available.
|
||||
|
||||
|
|
@ -512,54 +465,41 @@ class CampaignCallDispatcher:
|
|||
campaign: The campaign object
|
||||
timeout: Maximum time to wait for a slot (default 10 minutes)
|
||||
|
||||
Returns the slot_id which must be released when the call completes.
|
||||
Returns the slot which must be released when the call completes.
|
||||
|
||||
Raises:
|
||||
ConcurrentSlotAcquisitionError: If slot cannot be acquired within timeout
|
||||
"""
|
||||
# Get concurrent limit for organization
|
||||
org_concurrent_limit = await self.get_org_concurrent_limit(organization_id)
|
||||
|
||||
# Check for campaign-level max_concurrency in orchestrator_metadata
|
||||
# Check for campaign-level max_concurrency in orchestrator_metadata.
|
||||
# It caps this campaign's own concurrent calls via a campaign-scoped
|
||||
# counter — the org-wide limit still applies on top, but calls from
|
||||
# other sources (WebRTC, inbound, other campaigns) don't count
|
||||
# against the campaign's cap.
|
||||
campaign_max_concurrency = None
|
||||
if campaign.orchestrator_metadata:
|
||||
campaign_max_concurrency = campaign.orchestrator_metadata.get(
|
||||
"max_concurrency"
|
||||
)
|
||||
|
||||
# Use the lower of campaign limit and org limit
|
||||
if campaign_max_concurrency is not None:
|
||||
max_concurrent = min(campaign_max_concurrency, org_concurrent_limit)
|
||||
else:
|
||||
max_concurrent = org_concurrent_limit
|
||||
|
||||
# Track wait time for alerting
|
||||
wait_start = time.time()
|
||||
|
||||
# Wait until we can acquire a concurrent slot
|
||||
while True:
|
||||
slot_id = await rate_limiter.try_acquire_concurrent_slot(
|
||||
organization_id, max_concurrent
|
||||
try:
|
||||
return await call_concurrency.acquire_org_slot(
|
||||
organization_id,
|
||||
source=f"campaign:{campaign.id}",
|
||||
timeout=timeout,
|
||||
scope_key=(
|
||||
f"campaign:{campaign.id}"
|
||||
if campaign_max_concurrency is not None
|
||||
else None
|
||||
),
|
||||
scope_max_concurrent=campaign_max_concurrency,
|
||||
retry_interval=1,
|
||||
)
|
||||
if slot_id:
|
||||
return slot_id
|
||||
|
||||
# Check if we've been waiting too long
|
||||
wait_time = time.time() - wait_start
|
||||
if wait_time > timeout:
|
||||
raise ConcurrentSlotAcquisitionError(
|
||||
organization_id=organization_id,
|
||||
campaign_id=campaign.id,
|
||||
wait_time=wait_time,
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"Attempting to get a slot for {organization_id} {campaign.id}, "
|
||||
f"waited {wait_time:.1f}s"
|
||||
)
|
||||
|
||||
# Wait before retrying
|
||||
await asyncio.sleep(1)
|
||||
except CallConcurrencyLimitError as e:
|
||||
raise ConcurrentSlotAcquisitionError(
|
||||
organization_id=organization_id,
|
||||
campaign_id=campaign.id,
|
||||
wait_time=e.wait_time,
|
||||
) from e
|
||||
|
||||
async def acquire_from_number(
|
||||
self,
|
||||
|
|
@ -604,17 +544,9 @@ class CampaignCallDispatcher:
|
|||
Release concurrent slot and from_number when a call completes.
|
||||
Called by Twilio webhooks or workflow completion handlers.
|
||||
"""
|
||||
slot_released = False
|
||||
mapping = await rate_limiter.get_workflow_slot_mapping(workflow_run_id)
|
||||
if mapping:
|
||||
org_id, slot_id = mapping
|
||||
success = await rate_limiter.release_concurrent_slot(org_id, slot_id)
|
||||
if success:
|
||||
await rate_limiter.delete_workflow_slot_mapping(workflow_run_id)
|
||||
logger.info(
|
||||
f"Released concurrent slot for workflow run {workflow_run_id}"
|
||||
)
|
||||
slot_released = True
|
||||
slot_released = await call_concurrency.release_workflow_run_slot(
|
||||
workflow_run_id
|
||||
)
|
||||
|
||||
# Release from_number back to its (org, telephony config) pool
|
||||
from_number_mapping = await rate_limiter.get_workflow_from_number_mapping(
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
|
|
@ -8,6 +9,12 @@ from loguru import logger
|
|||
from api.constants import REDIS_URL
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ConcurrentSlotAcquisition:
|
||||
slot_id: str
|
||||
active_count: int
|
||||
|
||||
|
||||
class RateLimiter:
|
||||
"""Sliding window rate limiter to enforce strict per-second limits and concurrent call limits"""
|
||||
|
||||
|
|
@ -100,34 +107,71 @@ class RateLimiter:
|
|||
Try to acquire a concurrent call slot.
|
||||
Returns a unique slot_id if successful, None if limit reached.
|
||||
"""
|
||||
acquisition = await self.try_acquire_concurrent_slot_details(
|
||||
organization_id, max_concurrent
|
||||
)
|
||||
return acquisition.slot_id if acquisition else None
|
||||
|
||||
async def try_acquire_concurrent_slot_details(
|
||||
self,
|
||||
organization_id: int,
|
||||
max_concurrent: int = 20,
|
||||
*,
|
||||
scope_key: str | None = None,
|
||||
scope_max_concurrent: int | None = None,
|
||||
) -> Optional[ConcurrentSlotAcquisition]:
|
||||
"""
|
||||
Try to acquire a concurrent call slot.
|
||||
Returns the slot_id and post-acquire active count if successful,
|
||||
or None if the limit is reached.
|
||||
|
||||
When ``scope_key``/``scope_max_concurrent`` are provided, the slot is
|
||||
also registered in a secondary counter (``concurrent_calls:<scope_key>``,
|
||||
e.g. ``campaign:<id>``) and acquisition additionally requires that
|
||||
counter to be below ``scope_max_concurrent``. Both counters are
|
||||
updated atomically. The scope-scoped slot must be released with the
|
||||
same ``scope_key``.
|
||||
"""
|
||||
redis_client = await self._get_redis()
|
||||
|
||||
concurrent_key = f"concurrent_calls:{organization_id}"
|
||||
scope_concurrent_key = f"concurrent_calls:{scope_key}" if scope_key else ""
|
||||
now = time.time()
|
||||
stale_cutoff = now - self.stale_call_timeout
|
||||
|
||||
# Lua script for atomic operation
|
||||
# Lua script for atomic operation across the org counter and the
|
||||
# optional scope counter (empty scope key = org-only acquisition).
|
||||
lua_script = """
|
||||
local key = KEYS[1]
|
||||
local scope_key = KEYS[2]
|
||||
local now = tonumber(ARGV[1])
|
||||
local max_concurrent = tonumber(ARGV[2])
|
||||
local stale_cutoff = tonumber(ARGV[3])
|
||||
local slot_id = ARGV[4]
|
||||
|
||||
-- Remove stale entries (older than 30 minutes)
|
||||
local scope_max_concurrent = tonumber(ARGV[5])
|
||||
|
||||
-- Remove stale entries (older than the stale-call timeout)
|
||||
redis.call('ZREMRANGEBYSCORE', key, 0, stale_cutoff)
|
||||
|
||||
|
||||
-- Get current count
|
||||
local current_count = redis.call('ZCARD', key)
|
||||
|
||||
if current_count < max_concurrent then
|
||||
-- Add new slot
|
||||
redis.call('ZADD', key, now, slot_id)
|
||||
redis.call('EXPIRE', key, 3600) -- Expire after 1 hour
|
||||
return slot_id
|
||||
else
|
||||
|
||||
if current_count >= max_concurrent then
|
||||
return nil
|
||||
end
|
||||
|
||||
if scope_key ~= '' then
|
||||
redis.call('ZREMRANGEBYSCORE', scope_key, 0, stale_cutoff)
|
||||
if redis.call('ZCARD', scope_key) >= scope_max_concurrent then
|
||||
return nil
|
||||
end
|
||||
redis.call('ZADD', scope_key, now, slot_id)
|
||||
redis.call('EXPIRE', scope_key, 3600)
|
||||
end
|
||||
|
||||
redis.call('ZADD', key, now, slot_id)
|
||||
redis.call('EXPIRE', key, 3600) -- Expire after 1 hour
|
||||
return {slot_id, current_count + 1}
|
||||
"""
|
||||
|
||||
# Generate unique slot ID (timestamp + random component)
|
||||
|
|
@ -136,22 +180,38 @@ class RateLimiter:
|
|||
try:
|
||||
result = await redis_client.eval(
|
||||
lua_script,
|
||||
1,
|
||||
2,
|
||||
concurrent_key,
|
||||
scope_concurrent_key,
|
||||
now,
|
||||
max_concurrent,
|
||||
stale_cutoff,
|
||||
slot_id,
|
||||
scope_max_concurrent if scope_max_concurrent is not None else 0,
|
||||
)
|
||||
if not result:
|
||||
return None
|
||||
|
||||
acquired_slot_id, active_count = result
|
||||
return ConcurrentSlotAcquisition(
|
||||
slot_id=str(acquired_slot_id),
|
||||
active_count=int(active_count),
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.error(f"Concurrent limiter error: {e}")
|
||||
return None
|
||||
|
||||
async def release_concurrent_slot(self, organization_id: int, slot_id: str) -> bool:
|
||||
async def release_concurrent_slot(
|
||||
self,
|
||||
organization_id: int,
|
||||
slot_id: str,
|
||||
scope_key: str | None = None,
|
||||
) -> bool | None:
|
||||
"""
|
||||
Release a concurrent call slot.
|
||||
Returns True if slot was released, False otherwise.
|
||||
Release a concurrent call slot (and its scope counter entry, if any).
|
||||
Returns True if the slot was released, False if it was already gone
|
||||
(released/stale-expired), or None on a Redis error — callers that
|
||||
track cleanup state should keep it around for retry when None.
|
||||
"""
|
||||
if not slot_id:
|
||||
return False
|
||||
|
|
@ -161,6 +221,8 @@ class RateLimiter:
|
|||
|
||||
try:
|
||||
removed = await redis_client.zrem(concurrent_key, slot_id)
|
||||
if scope_key:
|
||||
await redis_client.zrem(f"concurrent_calls:{scope_key}", slot_id)
|
||||
if removed:
|
||||
logger.debug(
|
||||
f"Released concurrent slot {slot_id} for org {organization_id}"
|
||||
|
|
@ -168,7 +230,7 @@ class RateLimiter:
|
|||
return bool(removed)
|
||||
except Exception as e:
|
||||
logger.error(f"Error releasing concurrent slot: {e}")
|
||||
return False
|
||||
return None
|
||||
|
||||
async def get_concurrent_count(self, organization_id: int) -> int:
|
||||
"""
|
||||
|
|
@ -212,12 +274,62 @@ class RateLimiter:
|
|||
logger.error(f"Error storing workflow slot mapping: {e}")
|
||||
return False
|
||||
|
||||
async def store_workflow_slot_mapping_if_absent(
|
||||
self,
|
||||
workflow_run_id: int,
|
||||
organization_id: int,
|
||||
slot_id: str,
|
||||
scope_key: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Store the workflow_run_id -> concurrent slot mapping only if no mapping
|
||||
already exists. This prevents duplicate public/WebRTC starts for the
|
||||
same workflow run from overwriting the cleanup pointer.
|
||||
"""
|
||||
redis_client = await self._get_redis()
|
||||
mapping_key = f"workflow_slot_mapping:{workflow_run_id}"
|
||||
|
||||
lua_script = """
|
||||
local key = KEYS[1]
|
||||
local org_id = ARGV[1]
|
||||
local slot_id = ARGV[2]
|
||||
local ttl = tonumber(ARGV[3])
|
||||
local scope_key = ARGV[4]
|
||||
|
||||
if redis.call('EXISTS', key) == 1 then
|
||||
return 0
|
||||
end
|
||||
|
||||
redis.call('HSET', key, 'org_id', org_id, 'slot_id', slot_id)
|
||||
if scope_key ~= '' then
|
||||
redis.call('HSET', key, 'scope_key', scope_key)
|
||||
end
|
||||
redis.call('EXPIRE', key, ttl)
|
||||
return 1
|
||||
"""
|
||||
|
||||
try:
|
||||
stored = await redis_client.eval(
|
||||
lua_script,
|
||||
1,
|
||||
mapping_key,
|
||||
organization_id,
|
||||
slot_id,
|
||||
self.stale_call_timeout,
|
||||
scope_key or "",
|
||||
)
|
||||
return bool(stored)
|
||||
except Exception as e:
|
||||
logger.error(f"Error storing workflow slot mapping if absent: {e}")
|
||||
return False
|
||||
|
||||
async def get_workflow_slot_mapping(
|
||||
self, workflow_run_id: int
|
||||
) -> Optional[tuple[int, str]]:
|
||||
) -> Optional[tuple[int, str, str | None]]:
|
||||
"""
|
||||
Get the concurrent slot mapping for a workflow run.
|
||||
Returns (organization_id, slot_id) tuple or None if not found.
|
||||
Returns (organization_id, slot_id, scope_key) or None if not found;
|
||||
scope_key is None for slots acquired without a scope counter.
|
||||
"""
|
||||
redis_client = await self._get_redis()
|
||||
mapping_key = f"workflow_slot_mapping:{workflow_run_id}"
|
||||
|
|
@ -225,7 +337,11 @@ class RateLimiter:
|
|||
try:
|
||||
mapping = await redis_client.hgetall(mapping_key)
|
||||
if mapping and "org_id" in mapping and "slot_id" in mapping:
|
||||
return (int(mapping["org_id"]), mapping["slot_id"])
|
||||
return (
|
||||
int(mapping["org_id"]),
|
||||
mapping["slot_id"],
|
||||
mapping.get("scope_key") or None,
|
||||
)
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.error(f"Error getting workflow slot mapping: {e}")
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue