Coverage for agentos/tools/di_container.py: 0%

100 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 23:53 +0800

1""" 

2DIContainer — lightweight dependency injection container. 

3 

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""" 

12 

13from __future__ import annotations 

14 

15import inspect 

16from collections.abc import Callable 

17from enum import Enum 

18from threading import RLock 

19from typing import Any, TypeVar, get_type_hints 

20 

21T = TypeVar("T") 

22 

23 

24# ============================================================================ 

25# Lifetime 

26# ============================================================================ 

27 

28 

29class Lifetime(Enum): 

30 SINGLETON = "singleton" 

31 TRANSIENT = "transient" 

32 

33 

34# ============================================================================ 

35# Registration 

36# ============================================================================ 

37 

38 

39class Registration: 

40 __slots__ = ("interface", "implementation", "lifetime", "factory", "instance", "instance_lock") 

41 

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() 

56 

57 

58# ============================================================================ 

59# Circular Dependency Error 

60# ============================================================================ 

61 

62 

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)}") 

67 

68 

69# ============================================================================ 

70# DIContainer 

71# ============================================================================ 

72 

73 

74class DIContainer: 

75 """Lightweight dependency injection container. 

76 

77 Usage: 

78 container = DIContainer() 

79 

80 # Register singleton 

81 container.register(AbstractDB, ConcretePostgres, Lifetime.SINGLETON) 

82 

83 # Register transient 

84 container.register(AbstractCache, RedisCache, Lifetime.TRANSIENT) 

85 

86 # Register instance 

87 container.register_instance(ConfigService, config_obj) 

88 

89 # Resolve 

90 db = container.resolve(AbstractDB) 

91 

92 # Factory 

93 container.register_factory(AbstractQueue, lambda: build_queue()) 

94 """ 

95 

96 def __init__(self, parent: DIContainer | None = None): 

97 self._registrations: dict[Any, Registration] = {} 

98 self._lock = RLock() 

99 self._parent = parent 

100 

101 # ---------- register ---------- 

102 

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 ) 

118 

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 ) 

128 

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 ) 

140 

141 # ---------- resolve ---------- 

142 

143 def resolve(self, interface: type[T]) -> T: 

144 """Resolve and return an instance of the given interface.""" 

145 return self._resolve(interface, set()) 

146 

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]) 

151 

152 reg = self._get_registration(interface) 

153 resolving.add(interface) 

154 

155 try: 

156 # Instance already cached 

157 if reg.instance is not None: 

158 return reg.instance 

159 

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 

168 

169 # Transient: create every time 

170 return self._build(reg, resolving) 

171 finally: 

172 resolving.discard(interface) 

173 

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__}") 

181 

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() 

186 

187 # Constructor injection 

188 impl = reg.implementation or reg.interface 

189 hints = self._safe_get_type_hints(impl.__init__) 

190 kwargs: dict[str, Any] = {} 

191 

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 

206 

207 return impl(**kwargs) 

208 

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 {} 

215 

216 # ---------- scoped ---------- 

217 

218 def create_scope(self) -> DIContainer: 

219 """Create a scoped child container (snapshot of current registrations).""" 

220 return DIContainer(parent=self) 

221 

222 # ---------- check ---------- 

223 

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