Coverage for agentos/tools/di_container.py: 0%
100 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 13:14 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 13:14 +0800
1"""
2DIContainer — lightweight dependency injection container.
4Supports:
5 - Singleton and transient lifetimes
6 - Constructor autowiring via type annotations
7 - Factory registration
8 - Instance registration
9 - Scoped sub-containers (snapshot-based)
10 - Circular dependency detection
11"""
13from __future__ import annotations
15import inspect
16from collections.abc import Callable
17from enum import Enum
18from threading import RLock
19from typing import Any, TypeVar, get_type_hints
21T = TypeVar("T")
24# ============================================================================
25# Lifetime
26# ============================================================================
29class Lifetime(Enum):
30 SINGLETON = "singleton"
31 TRANSIENT = "transient"
34# ============================================================================
35# Registration
36# ============================================================================
39class Registration:
40 __slots__ = ("interface", "implementation", "lifetime", "factory", "instance", "instance_lock")
42 def __init__(
43 self,
44 interface: type,
45 implementation: type | None = None,
46 lifetime: Lifetime = Lifetime.TRANSIENT,
47 factory: Callable[[], Any] | None = None,
48 instance: Any = None,
49 ):
50 self.interface = interface
51 self.implementation = implementation
52 self.lifetime = lifetime
53 self.factory = factory
54 self.instance = instance
55 self.instance_lock = RLock()
58# ============================================================================
59# Circular Dependency Error
60# ============================================================================
63class CircularDependencyError(Exception):
64 def __init__(self, chain: list):
65 self.chain = chain
66 super().__init__(f"Circular dependency detected: {' → '.join(str(c) for c in chain)}")
69# ============================================================================
70# DIContainer
71# ============================================================================
74class DIContainer:
75 """Lightweight dependency injection container.
77 Usage:
78 container = DIContainer()
80 # Register singleton
81 container.register(AbstractDB, ConcretePostgres, Lifetime.SINGLETON)
83 # Register transient
84 container.register(AbstractCache, RedisCache, Lifetime.TRANSIENT)
86 # Register instance
87 container.register_instance(ConfigService, config_obj)
89 # Resolve
90 db = container.resolve(AbstractDB)
92 # Factory
93 container.register_factory(AbstractQueue, lambda: build_queue())
94 """
96 def __init__(self, parent: DIContainer | None = None):
97 self._registrations: dict[Any, Registration] = {}
98 self._lock = RLock()
99 self._parent = parent
101 # ---------- register ----------
103 def register(
104 self,
105 interface: type,
106 implementation: type | None = None,
107 lifetime: Lifetime = Lifetime.TRANSIENT,
108 ) -> None:
109 """Register an interface with its implementation."""
110 if implementation is None:
111 implementation = interface
112 with self._lock:
113 self._registrations[interface] = Registration(
114 interface=interface,
115 implementation=implementation,
116 lifetime=lifetime,
117 )
119 def register_instance(self, interface: type, instance: Any) -> None:
120 """Register a pre-built instance."""
121 with self._lock:
122 self._registrations[interface] = Registration(
123 interface=interface,
124 implementation=type(instance),
125 lifetime=Lifetime.SINGLETON,
126 instance=instance,
127 )
129 def register_factory(
130 self, interface: type, factory: Callable[[], Any], lifetime: Lifetime = Lifetime.TRANSIENT
131 ) -> None:
132 """Register a factory callable for the interface."""
133 with self._lock:
134 self._registrations[interface] = Registration(
135 interface=interface,
136 implementation=None,
137 lifetime=lifetime,
138 factory=factory,
139 )
141 # ---------- resolve ----------
143 def resolve(self, interface: type[T]) -> T:
144 """Resolve and return an instance of the given interface."""
145 return self._resolve(interface, set())
147 def _resolve(self, interface: type, resolving: set[type]) -> Any:
148 # Check circular deps
149 if interface in resolving:
150 raise CircularDependencyError(list(resolving) + [interface])
152 reg = self._get_registration(interface)
153 resolving.add(interface)
155 try:
156 # Instance already cached
157 if reg.instance is not None:
158 return reg.instance
160 # Singleton: create once
161 if reg.lifetime == Lifetime.SINGLETON:
162 with reg.instance_lock:
163 if reg.instance is not None:
164 return reg.instance
165 instance = self._build(reg, resolving)
166 reg.instance = instance
167 return instance
169 # Transient: create every time
170 return self._build(reg, resolving)
171 finally:
172 resolving.discard(interface)
174 def _get_registration(self, interface: type) -> Registration:
175 with self._lock:
176 if interface in self._registrations:
177 return self._registrations[interface]
178 if self._parent:
179 return self._parent._get_registration(interface)
180 raise KeyError(f"No registration for {interface.__name__}")
182 def _build(self, reg: Registration, resolving: set[type]) -> Any:
183 # Factory takes priority
184 if reg.factory is not None:
185 return reg.factory()
187 # Constructor injection
188 impl = reg.implementation or reg.interface
189 hints = self._safe_get_type_hints(impl.__init__)
190 kwargs: dict[str, Any] = {}
192 for param_name, param in inspect.signature(impl.__init__).parameters.items():
193 if param_name == "self":
194 continue
195 param_type = hints.get(param_name)
196 if param_type is not None:
197 try:
198 kwargs[param_name] = self._resolve(param_type, resolving.copy())
199 except (KeyError, CircularDependencyError):
200 if param.default is not inspect.Parameter.empty:
201 kwargs[param_name] = param.default
202 else:
203 raise
204 elif param.default is not inspect.Parameter.empty:
205 kwargs[param_name] = param.default
207 return impl(**kwargs)
209 @staticmethod
210 def _safe_get_type_hints(func) -> dict[str, Any]:
211 try:
212 return get_type_hints(func)
213 except Exception:
214 return {}
216 # ---------- scoped ----------
218 def create_scope(self) -> DIContainer:
219 """Create a scoped child container (snapshot of current registrations)."""
220 return DIContainer(parent=self)
222 # ---------- check ----------
224 def is_registered(self, interface: type) -> bool:
225 with self._lock:
226 if interface in self._registrations:
227 return True
228 if self._parent:
229 return self._parent.is_registered(interface)
230 return False