diff --git a/crates/fluxqueue-worker/src/task.rs b/crates/fluxqueue-worker/src/task.rs index 0907f42..9ad458a 100644 --- a/crates/fluxqueue-worker/src/task.rs +++ b/crates/fluxqueue-worker/src/task.rs @@ -381,7 +381,7 @@ fn get_registry(module_path: &str, queue_name: &str) -> Result ) .collect::, _>>()?; - let contexts: HashMap, Arc>> = registry + let mut contexts: HashMap, Arc>> = registry .get_item("contexts")? .expect("contexts missing") .cast::() @@ -389,11 +389,14 @@ fn get_registry(module_path: &str, queue_name: &str) -> Result .iter() .filter_map(|(key, value): (Bound, Bound)| { let name: String = key.extract().ok()?; - let func: Py = value.unbind(); - Some((Arc::new(name), Arc::new(func))) + let class: Py = value.unbind(); + Some((Arc::new(name), Arc::new(class))) }) .collect(); + let (base_name, base_class) = get_base_context_class(py)?; + contexts.insert(base_name, base_class); + Ok((tasks, contexts)) })?; @@ -416,6 +419,15 @@ fn get_task_metadata(py: Python<'_>, task: Arc) -> Result> { Ok(task_metadata) } +fn get_base_context_class(py: Python<'_>) -> Result<(Arc, Arc>)> { + let module = py.import("fluxqueue.context")?; + let context_class = module.getattr("Context")?; + Ok(( + Arc::new("_Context".to_string()), + Arc::new(context_class.unbind()), + )) +} + fn is_coroutine(py: Python<'_>, func: Arc>) -> Result { let inspect = py.import("inspect")?; let is_coro: bool = inspect diff --git a/python/fluxqueue/context.py b/python/fluxqueue/context.py index 6296ef7..9ce638c 100644 --- a/python/fluxqueue/context.py +++ b/python/fluxqueue/context.py @@ -29,7 +29,7 @@ class Context: to provide domain-specific resources. """ - __fluxqueue_context__: str | None = None + __fluxqueue_context__: str | None = "_Context" def __init__(self) -> None: self._thread_local = threading.local() @@ -73,7 +73,10 @@ async def _run_async_task( self._metadata_var.reset(token) def __init_subclass__(cls) -> None: - if not cls.__fluxqueue_context__: + if cls.__name__ == "_Context": + raise ValueError("Context subclass cannot be named '_Context'") + + if not cls.__fluxqueue_context__ or cls.__fluxqueue_context__ == "_Context": cls.__fluxqueue_context__ = cls.__name__