Skip to content

Commit 2636501

Browse files
committed
feat: use infrastructure class as global wrapper, add meter event handler, initial integration of nudging event
1 parent 3670f83 commit 2636501

17 files changed

Lines changed: 313 additions & 217 deletions

File tree

config/clients.yaml

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,13 @@
55
clients:
66
dataset_api:
77
class: celine.dt.core.clients.dataset_api:DatasetSqlApiClient
8+
scope: dataset.query
89
config:
910
base_url: "${DATASET_API_BASE_URL:-http://api.celine.localhost/datasets}"
1011
timeout: 30.0
12+
rec_registry_admin:
13+
class: celine.sdk.rec_registry:RecRegistryAdminClient
14+
scope: "rec-registry.lookup"
15+
config:
16+
base_url: "${REC_REGISTRY_URL:-http://api.celine.localhost/rec-registry}"
17+
timeout: 5.0

src/celine/dt/contracts/__init__.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,8 +3,10 @@
33
from celine.dt.contracts.events import DTEvent, EventSource, EventSeverity
44
from celine.dt.contracts.component import DTComponent
55
from celine.dt.contracts.simulation import DTSimulation, SimulationDescriptor
6-
from celine.dt.contracts.subscription import SubscriptionSpec, EventHandler, EventContext
6+
from celine.dt.contracts.subscription import SubscriptionSpec, EventHandler, EventContext, RouteDef
77
from celine.dt.contracts.values import ValueFetcherSpec
8+
from celine.dt.contracts.infrastructure import Infrastructure
9+
from celine.dt.contracts.app import AppState
810

911
# Broker types re-exported from SDK
1012
from celine.sdk.broker import (
@@ -21,7 +23,9 @@
2123
"DTEvent", "EventSource", "EventSeverity",
2224
"DTComponent",
2325
"DTSimulation", "SimulationDescriptor",
24-
"SubscriptionSpec", "EventHandler", "EventContext",
26+
"SubscriptionSpec", "EventHandler", "EventContext", "RouteDef",
2527
"ValueFetcherSpec",
2628
"Broker", "BrokerMessage", "MqttBroker", "MqttConfig", "PublishResult", "QoS",
29+
"Infrastructure",
30+
"AppState"
2731
]

src/celine/dt/contracts/app.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,10 @@
1+
from dataclasses import dataclass
2+
3+
from starlette.datastructures import State
4+
from celine.dt.contracts.infrastructure import Infrastructure
5+
from celine.sdk.auth import TokenProvider
6+
7+
@dataclass
8+
class AppState(State):
9+
infra: Infrastructure
10+
token_provider: TokenProvider | None = None
Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,54 @@
1+
# celine/dt/core/infrastructure.py
2+
from __future__ import annotations
3+
from dataclasses import dataclass, field, replace
4+
from typing import TYPE_CHECKING, Any, Optional, TypeVar
5+
6+
from celine.dt.core.broker.service import BrokerService
7+
from celine.dt.core.clients.registry import ClientsRegistry
8+
from celine.dt.core.values.service import ValuesService, ValuesRegistry
9+
from celine.dt.core.simulation.registry import SimulationRegistry
10+
11+
if TYPE_CHECKING:
12+
from celine.dt.core.domain.registry import DomainRegistry
13+
from celine.dt.core.broker.subscriptions import SubscriptionManager
14+
from celine.sdk.auth import TokenProvider
15+
16+
17+
@dataclass
18+
class Infrastructure:
19+
20+
# always present after create_app
21+
broker: BrokerService
22+
values_service: ValuesService
23+
values_registry: ValuesRegistry
24+
clients_registry: ClientsRegistry
25+
simulation_registry: SimulationRegistry
26+
27+
# set after lifespan finalization
28+
_domain_registry: Optional[DomainRegistry] = field(default=None)
29+
_subscription_manager: Optional[SubscriptionManager] = field(default=None)
30+
_token_provider: Optional[TokenProvider] = field(default=None)
31+
32+
overrides: dict[str, Any] = field(default_factory=dict)
33+
34+
@property
35+
def domain_registry(self) -> DomainRegistry:
36+
if self._domain_registry is None:
37+
raise RuntimeError("Infrastructure.domains not set yet - domain loading incomplete")
38+
return self._domain_registry
39+
40+
@property
41+
def subscription_manager(self) -> SubscriptionManager:
42+
if self._subscription_manager is None:
43+
raise RuntimeError("Infrastructure.subscription_manager not set yet - lifespan incomplete")
44+
return self._subscription_manager
45+
46+
@property
47+
def token_provider(self) -> TokenProvider:
48+
if self._token_provider is None:
49+
raise RuntimeError("Infrastructure.token_provider not set yet - lifespan incomplete")
50+
return self._token_provider
51+
52+
def with_overrides(self, overrides: dict[str, Any]) -> Infrastructure:
53+
"""Return a shallow copy with domain-specific overrides applied."""
54+
return replace(self, overrides=overrides)

src/celine/dt/contracts/routes.py

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -90,9 +90,5 @@ def from_dataclass(obj: Any) -> "FetchResultSchema":
9090
return FetchResultSchema.model_validate(obj.__dict__)
9191

9292

93-
class DescribeResponseSchema(BaseModel):
94-
payload: GenericPayload
95-
96-
9793
class SummaryResponseSchema(BaseModel):
9894
payload: GenericPayload

src/celine/dt/contracts/subscription.py

Lines changed: 7 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -12,12 +12,15 @@
1212
from typing import TYPE_CHECKING, Any, Awaitable, Callable, Protocol, runtime_checkable
1313
from uuid import uuid4
1414

15-
from celine.dt.contracts.events import DTEvent
15+
16+
from celine.dt.contracts import DTEvent
1617

1718
if TYPE_CHECKING:
19+
from celine.dt.core.clients.registry import ClientsRegistry
20+
from celine.dt.core.domain.registry import DomainRegistry
21+
from celine.dt.contracts import Infrastructure
1822
from celine.dt.core.broker.service import BrokerService
1923
from celine.dt.core.domain.base import DTDomain
20-
from celine.dt.core.domain.registry import DomainRegistry
2124
from celine.dt.core.values.service import ValuesService
2225

2326

@@ -32,35 +35,16 @@ async def on_run_completed(event: DTEvent, ctx: EventContext) -> None:
3235
users = await ctx.values.fetch("pipeline-reactor.affected_users", {...})
3336
await ctx.broker.publish_event(topic=f"celine/nudging/{user_id}", payload={...})
3437
"""
35-
3638
topic: str
3739
broker_name: str
3840
received_at: datetime
39-
40-
# App-scope infrastructure — always populated by SubscriptionManager
41-
broker: BrokerService
42-
values: ValuesService
43-
44-
# Domain registry — for get_dt() lookups
45-
registry: DomainRegistry | None = None
46-
41+
infra: Infrastructure
4742
entity_id: str | None = None
4843
message_id: str | None = None
4944
raw_payload: bytes | None = None
5045

5146
def get_dt(self, domain_type: str) -> DTDomain:
52-
"""Return the DTDomain instance for the given domain_type.
53-
54-
Usage in a plain @on_event handler::
55-
56-
@on_event("pipeline.run.completed", topics=["celine/pipelines/runs/+"])
57-
async def on_run_completed(event: DTEvent, ctx: EventContext) -> None:
58-
participant = ctx.get_dt("participant")
59-
entity = await participant.resolve_entity(participant_id)
60-
"""
61-
if self.registry is None:
62-
raise RuntimeError("EventContext has no registry — check SubscriptionManager setup")
63-
return self.registry.get_by_type(domain_type)
47+
return self.infra.domain_registry.get_by_type(domain_type)
6448

6549

6650
EventHandler = Callable[[DTEvent, EventContext], Awaitable[None]] # type: ignore[type-arg]

src/celine/dt/core/broker/subscriptions.py

Lines changed: 13 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -11,11 +11,7 @@
1111

1212
from celine.sdk.broker import QoS, ReceivedMessage, SubscribeResult
1313

14-
from celine.dt.contracts.events import DTEvent, EventSource
15-
from celine.dt.contracts.subscription import EventContext, EventHandler, RouteDef, SubscriptionSpec
16-
from celine.dt.core.broker.service import BrokerService
17-
from celine.dt.core.domain.registry import DomainRegistry
18-
from celine.dt.core.values.service import ValuesService
14+
from celine.dt.contracts import DTEvent, EventSource, EventContext, EventHandler, RouteDef, SubscriptionSpec, Infrastructure
1915

2016
logger = logging.getLogger(__name__)
2117

@@ -156,23 +152,23 @@ class SubscriptionManager:
156152
def __init__(
157153
self,
158154
*,
159-
broker_service: BrokerService,
160-
values_service: ValuesService,
161-
domain_registry: DomainRegistry | None = None,
155+
infra: Infrastructure,
162156
domains: list[Any] | None = None,
163157
handler_specs: list[SubscriptionSpec] | None = None,
164158
default_qos: QoS = QoS.AT_LEAST_ONCE,
165159
default_broker_name: str | None = None,
166160
) -> None:
167-
self._broker_service = broker_service
168-
self._values_service = values_service
169-
self._domain_registry = domain_registry
161+
self._infra = infra
170162
self._domains = domains or []
171163
self._handler_specs = handler_specs or []
172164
self._default_qos = default_qos
173165
self._default_broker_name = default_broker_name
174166
self._active: list[ActiveSubscription] = []
175167

168+
@property
169+
def infra(self) -> Infrastructure:
170+
return self._infra
171+
176172
async def start(self) -> None:
177173
# Domain instances
178174
for domain in self._domains:
@@ -208,7 +204,7 @@ async def _register_specs(
208204
broker_name=broker_name or "<default>",
209205
)
210206

211-
res: SubscribeResult = await self._broker_service.subscribe(
207+
res: SubscribeResult = await self.infra.broker.subscribe(
212208
topics=topics,
213209
handler=handler,
214210
broker_name=broker_name,
@@ -240,7 +236,7 @@ async def _register_specs(
240236
async def stop(self) -> None:
241237
for sub in list(self._active):
242238
try:
243-
await self._broker_service.unsubscribe(
239+
await self.infra.broker.unsubscribe(
244240
subscription_id=sub.subscription_id,
245241
broker_name=sub.broker_name,
246242
)
@@ -257,22 +253,20 @@ def _wrap_handler(
257253
spec: SubscriptionSpec,
258254
broker_name: str,
259255
) -> Callable[[ReceivedMessage], Awaitable[None]]:
260-
broker_service = self._broker_service
261-
values_service = self._values_service
262-
domain_registry = self._domain_registry
256+
broker_service = self.infra.broker
257+
values_service = self.infra.values_service
258+
domain_registry = self.infra.domain_registry
263259

264260
async def _handler(msg: ReceivedMessage) -> None:
265261
try:
266262
event = _dt_event_from_received(
267263
source_name=source_name, spec=spec, msg=msg
268264
)
269265
ctx = EventContext(
266+
infra=self.infra,
270267
topic=msg.topic,
271268
broker_name=broker_name,
272269
received_at=msg.timestamp or datetime.now(timezone.utc),
273-
broker=broker_service,
274-
values=values_service,
275-
registry=domain_registry,
276270
entity_id=spec.metadata.get("entity_id"),
277271
message_id=msg.message_id,
278272
raw_payload=msg.raw_payload,

src/celine/dt/core/clients/dataset_api.py

Lines changed: 17 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -5,11 +5,12 @@
55
from __future__ import annotations
66

77
import logging
8-
from typing import TYPE_CHECKING, Any, AsyncIterator, Optional
8+
from typing import TYPE_CHECKING, Any, AsyncIterator
99
import httpx
1010

1111
if TYPE_CHECKING:
1212
from celine.dt.api.context import Ctx
13+
from celine.sdk.auth.provider import TokenProvider
1314

1415
logger = logging.getLogger(__name__)
1516

@@ -28,28 +29,33 @@ def __init__(
2829
*,
2930
base_url: str,
3031
timeout: float = 30.0,
31-
token_provider: Any | None = None,
32+
token_provider: TokenProvider | None = None,
3233
) -> None:
3334
self._base = base_url.rstrip("/")
3435
self._timeout = timeout
3536
self._token_provider = token_provider
3637

37-
async def _headers(self) -> dict[str, str]:
38-
if not self._token_provider:
39-
return {}
40-
token = await self._token_provider.get_token()
41-
return {"Authorization": f"Bearer {token.access_token}"}
38+
async def _headers(self, user_token: str | None = None) -> dict[str, str]:
39+
40+
token = user_token
41+
if not user_token:
42+
if not self._token_provider:
43+
return {}
44+
client_token = await self._token_provider.get_token()
45+
token = client_token.access_token
46+
47+
return {"Authorization": f"Bearer {token}"}
4248

4349
async def query(
4450
self, *, sql: str, limit: int = 1000, offset: int = 0, ctx: Ctx | None = None
4551
) -> list[dict[str, Any]]:
46-
headers = await self._headers()
4752

53+
token: str | None = None
4854
if ctx and ctx.token:
4955
token = ctx.token
50-
headers["Authorization"] = (
51-
token if token.lower().startswith("bearer") else f"Bearer {token}"
52-
)
56+
token = token if token.lower().startswith("bearer") else f"Bearer {token}"
57+
58+
headers = await self._headers(token)
5359

5460
async with httpx.AsyncClient(timeout=self._timeout) as client:
5561
try:

src/celine/dt/core/clients/loader.py

Lines changed: 19 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -8,17 +8,21 @@
88
import logging
99
from typing import Any, Iterable
1010

11+
from celine.dt.core.config import settings
1112
from celine.dt.core.clients.registry import ClientsRegistry
1213
from celine.dt.core.loader import import_attr, load_yaml_files, substitute_env_vars
14+
from celine.sdk.auth.provider import TokenProvider
15+
from celine.sdk.auth import OidcClientCredentialsProvider
1316

1417
logger = logging.getLogger(__name__)
1518

1619

20+
1721
def load_and_register_clients(
1822
*,
1923
patterns: Iterable[str],
2024
registry: ClientsRegistry,
21-
token_provider: Any | None = None,
25+
token_provider: TokenProvider | None = None,
2226
) -> None:
2327
"""Load client definitions from YAML and register live instances.
2428
@@ -42,26 +46,26 @@ def load_and_register_clients(
4246
for data in yamls:
4347
for name, spec in (data.get("clients") or {}).items():
4448
class_path = spec.get("class")
45-
if not class_path:
46-
raise ValueError(f"Client '{name}' missing 'class' field")
47-
49+
scope = spec.get("scope")
4850
raw_config = substitute_env_vars(spec.get("config", {}))
4951

50-
logger.info("Loading client '%s' from '%s'", name, class_path)
51-
try:
52-
cls = import_attr(class_path)
53-
except (ImportError, AttributeError):
54-
logger.exception("Failed to import client class '%s'", class_path)
55-
raise
56-
52+
cls = import_attr(class_path)
5753
kwargs = dict(raw_config)
5854

59-
# Inject token_provider if the constructor accepts it
6055
sig = inspect.signature(cls.__init__)
6156
if "token_provider" in sig.parameters:
6257
kwargs["token_provider"] = token_provider
58+
if scope and isinstance(token_provider, OidcClientCredentialsProvider):
59+
if settings.oidc.client_id and settings.oidc.client_secret:
60+
kwargs["token_provider"] = OidcClientCredentialsProvider(
61+
base_url=settings.oidc.base_url,
62+
client_id=settings.oidc.client_id,
63+
client_secret=settings.oidc.client_secret,
64+
scope=scope,
65+
timeout=settings.oidc.timeout,
66+
)
67+
else:
68+
logger.warning(f"Cannot initialize OIDC token provider for aud={scope}: missing client_id / client_secret")
6369

64-
client = cls(**kwargs)
65-
registry.register(name, client)
66-
70+
registry.register(name, cls(**kwargs))
6771
logger.info("Registered %d client(s): %s", len(registry.list()), registry.list())

0 commit comments

Comments
 (0)