-
Notifications
You must be signed in to change notification settings - Fork 44
Expand file tree
/
Copy pathenqueue_states.py
More file actions
73 lines (62 loc) · 2.54 KB
/
Copy pathenqueue_states.py
File metadata and controls
73 lines (62 loc) · 2.54 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
import asyncio
import time
from ..models.enqueue_request import EnqueueRequestModel
from ..models.enqueue_response import EnqueueResponseModel, StateModel
from ..models.db.state import State
from ..models.state_status_enum import StateStatusEnum
from app.singletons.logs_manager import LogsManager
from pymongo import ReturnDocument
logger = LogsManager().get_logger()
async def find_state(namespace_name: str, nodes: list[str]) -> State | None:
current_time_ms = int(time.time() * 1000)
data = await State.get_pymongo_collection().find_one_and_update(
{
"namespace_name": namespace_name,
"status": StateStatusEnum.CREATED,
"node_name": {
"$in": nodes
},
"enqueue_after": {"$lte": current_time_ms}
},
{
"$set": {
"status": StateStatusEnum.QUEUED,
"queued_at": current_time_ms
}
},
return_document=ReturnDocument.AFTER
)
return State(**data) if data else None
async def enqueue_states(namespace_name: str, body: EnqueueRequestModel, x_exosphere_request_id: str) -> EnqueueResponseModel:
try:
logger.info(f"Enqueuing states for namespace {namespace_name}", x_exosphere_request_id=x_exosphere_request_id)
# Create tasks for parallel execution
tasks = [find_state(namespace_name, body.nodes) for _ in range(body.batch_size)]
results = await asyncio.gather(*tasks, return_exceptions=True)
# Filter out None results and exceptions
states = []
for result in results:
if isinstance(result, Exception):
logger.error(f"Error finding state: {result}", x_exosphere_request_id=x_exosphere_request_id)
continue
if result is not None:
states.append(result)
response = EnqueueResponseModel(
count=len(states),
namespace=namespace_name,
status=StateStatusEnum.QUEUED,
states=[
StateModel(
state_id=str(state.id),
node_name=state.node_name,
identifier=state.identifier,
inputs=state.inputs,
created_at=state.created_at
)
for state in states
]
)
return response
except Exception as e:
logger.error(f"Error enqueuing states for namespace {namespace_name}", x_exosphere_request_id=x_exosphere_request_id, error=e)
raise e