Skip to content

Commit e6de04e

Browse files
committed
adding logging to runtime
1 parent a813edd commit e6de04e

2 files changed

Lines changed: 65 additions & 8 deletions

File tree

python-sdk/exospherehost/runtime.py

Lines changed: 63 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,46 @@
11
import asyncio
22
import os
3+
import logging
4+
import traceback
5+
36
from asyncio import Queue, sleep
47
from typing import List, Dict
5-
68
from pydantic import BaseModel
79
from .node.BaseNode import BaseNode
810
from aiohttp import ClientSession
9-
from logging import getLogger
1011

11-
logger = getLogger(__name__)
12+
logger = logging.getLogger(__name__)
13+
14+
def _setup_default_logging():
15+
"""
16+
Setup default logging only if no handlers are configured.
17+
Respects user's existing logging configuration.
18+
"""
19+
root_logger = logging.getLogger()
20+
21+
# Don't interfere if user has already configured logging
22+
if root_logger.handlers:
23+
return
24+
25+
# Allow users to disable default logging
26+
if os.environ.get('EXOSPHERE_DISABLE_DEFAULT_LOGGING'):
27+
return
28+
29+
# Get log level from environment or default to INFO
30+
log_level_name = os.environ.get('EXOSPHERE_LOG_LEVEL', 'INFO').upper()
31+
log_level = getattr(logging, log_level_name, logging.INFO)
32+
33+
# Setup basic configuration with clean formatting
34+
logging.basicConfig(
35+
level=log_level,
36+
format='%(asctime)s | %(levelname)s | %(message)s',
37+
datefmt='%Y-%m-%d %H:%M:%S'
38+
)
39+
40+
# Log that we're using default configuration
41+
logger = logging.getLogger(__name__)
42+
logger.debug(f"ExosphereHost: Using default logging configuration (level: {log_level_name})")
43+
1244

1345
class Runtime:
1446
"""
@@ -48,6 +80,9 @@ class Runtime:
4880
"""
4981

5082
def __init__(self, namespace: str, name: str, nodes: List[type[BaseNode]], state_manager_uri: str | None = None, key: str | None = None, batch_size: int = 16, workers: int = 4, state_manage_version: str = "v0", poll_interval: int = 1):
83+
84+
_setup_default_logging()
85+
5186
self._name = name
5287
self._namespace = namespace
5388
self._key = key
@@ -72,8 +107,10 @@ def _set_config_from_env(self):
72107
Set configuration from environment variables if not provided.
73108
"""
74109
if self._state_manager_uri is None:
110+
logger.info("State manager URI not provided, using environment variable EXOSPHERE_STATE_MANAGER_URI")
75111
self._state_manager_uri = os.environ.get("EXOSPHERE_STATE_MANAGER_URI")
76112
if self._key is None:
113+
logger.info("API key not provided, using environment variable EXOSPHERE_API_KEY")
77114
self._key = os.environ.get("EXOSPHERE_API_KEY")
78115

79116
def _validate_runtime(self):
@@ -130,6 +167,7 @@ async def _register(self):
130167
Raises:
131168
RuntimeError: If registration fails.
132169
"""
170+
logger.info(f"Registering nodes: {[f"{self._namespace}/{node.__name__}" for node in self._nodes]}")
133171
async with ClientSession() as session:
134172
endpoint = self._get_register_endpoint()
135173
body = {
@@ -153,8 +191,10 @@ async def _register(self):
153191
res = await response.json()
154192

155193
if response.status != 200:
194+
logger.error(f"Failed to register nodes: {res}")
156195
raise RuntimeError(f"Failed to register nodes: {res}")
157196

197+
logger.info(f"Registered nodes: {[f"{self._namespace}/{node.__name__}" for node in self._nodes]}")
158198
return res
159199

160200
async def _enqueue_call(self):
@@ -174,6 +214,7 @@ async def _enqueue_call(self):
174214

175215
if response.status != 200:
176216
logger.error(f"Failed to enqueue states: {res}")
217+
raise RuntimeError(f"Failed to enqueue states: {res}")
177218

178219
return res
179220

@@ -189,8 +230,10 @@ async def _enqueue(self):
189230
data = await self._enqueue_call()
190231
for state in data.get("states", []):
191232
await self._state_queue.put(state)
233+
logger.info(f"Enqueued states: {len(data.get('states', []))}")
192234
except Exception as e:
193235
logger.error(f"Error enqueuing states: {e}")
236+
raise
194237

195238
await sleep(self._poll_interval)
196239

@@ -212,6 +255,8 @@ async def _notify_executed(self, state_id: str, outputs: List[BaseNode.Outputs])
212255

213256
if response.status != 200:
214257
logger.error(f"Failed to notify executed state {state_id}: {res}")
258+
259+
logger.info(f"Notified executed state {state_id} with outputs: {outputs} for node {self._node_mapping[state_id].__name__}")
215260

216261
async def _notify_errored(self, state_id: str, error: str):
217262
"""
@@ -232,6 +277,8 @@ async def _notify_errored(self, state_id: str, error: str):
232277
if response.status != 200:
233278
logger.error(f"Failed to notify errored state {state_id}: {res}")
234279

280+
logger.info(f"Notified errored state {state_id} with error: {error} for node {self._node_mapping[state_id].__name__}")
281+
235282
async def _get_secrets(self, state_id: str) -> Dict[str, str]:
236283
"""
237284
Get secrets for a state.
@@ -306,21 +353,28 @@ def _validate_nodes(self):
306353
if len(errors) > 0:
307354
raise ValueError("Following errors while validating nodes: " + "\n".join(errors))
308355

309-
async def _worker(self):
356+
async def _worker(self, idx: int):
310357
"""
311358
Worker task that processes states from the queue.
312359
313360
Continuously fetches states from the queue, executes the corresponding node,
314361
and notifies the state manager of the result.
315362
"""
363+
logger.info(f"Starting worker thread {idx} for nodes: {[f"{self._namespace}/{node.__name__}" for node in self._nodes]}")
364+
316365
while True:
317366
state = await self._state_queue.get()
318367

319368
try:
320369
node = self._node_mapping[state["node_name"]]
370+
logger.info(f"Executing state {state['state_id']} for node {node.__name__}")
371+
321372
secrets = await self._get_secrets(state["state_id"])
322-
outputs = await node()._execute(node.Inputs(**state["inputs"]), node.Secrets(**secrets["secrets"]))
373+
logger.info(f"Got secrets for state {state['state_id']} for node {node.__name__}")
323374

375+
outputs = await node()._execute(node.Inputs(**state["inputs"]), node.Secrets(**secrets["secrets"])) # type: ignore
376+
logger.info(f"Got outputs for state {state['state_id']} for node {node.__name__}")
377+
324378
if outputs is None:
325379
outputs = []
326380

@@ -330,6 +384,9 @@ async def _worker(self):
330384
await self._notify_executed(state["state_id"], outputs)
331385

332386
except Exception as e:
387+
logger.error(f"Error executing state {state['state_id']} for node {node.__name__}: {e}")
388+
logger.error(traceback.format_exc())
389+
333390
await self._notify_errored(state["state_id"], str(e))
334391

335392
self._state_queue.task_done() # type: ignore
@@ -346,7 +403,7 @@ async def _start(self):
346403
await self._register()
347404

348405
poller = asyncio.create_task(self._enqueue())
349-
worker_tasks = [asyncio.create_task(self._worker()) for _ in range(self._workers)]
406+
worker_tasks = [asyncio.create_task(self._worker(idx)) for idx in range(self._workers)]
350407

351408
await asyncio.gather(poller, *worker_tasks)
352409

python-sdk/tests/test_runtime_validation.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -15,11 +15,11 @@ class Secrets(BaseModel):
1515
api_key: str
1616

1717
async def execute(self):
18-
return self.Outputs(message=f"hi {self.inputs.name}")
18+
return self.Outputs(message=f"hi {self.inputs.name}") # type: ignore
1919

2020

2121
class BadNodeWrongInputsBase(BaseNode):
22-
Inputs = object # not a pydantic BaseModel
22+
Inputs = object # not a pydantic BaseModel # type: ignore
2323
class Outputs(BaseModel):
2424
message: str
2525
class Secrets(BaseModel):

0 commit comments

Comments
 (0)