Coverage for src/lexigram/auth/authn/ldap.py: 0%

189 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-25 12:26 +0800

1"""LDAP authentication manager implementation.""" 

2 

3from __future__ import annotations 

4 

5from dataclasses import dataclass 

6from typing import Any 

7 

8from lexigram.logging import get_logger 

9 

10try: 

11 import ldap3 

12 from ldap3 import ALL, BASE, SUBTREE, Connection, Server 

13 from ldap3.core.exceptions import LDAPException 

14 

15 LDAP_AVAILABLE = True 

16except ImportError: 

17 LDAP_AVAILABLE = False 

18 ldap3 = None 

19 # Define placeholder types for when LDAP is not available 

20 Connection = Any 

21 Server = Any 

22 

23from lexigram.auth.exceptions import AuthError 

24from lexigram.contracts.web import HTTPClientProtocol 

25from lexigram.validation import SecretStr 

26 

27logger = get_logger(__name__) 

28 

29 

30@dataclass 

31class LDAPProvider: 

32 """LDAP provider configuration.""" 

33 

34 name: str 

35 server_url: str 

36 bind_dn: str | None = None 

37 bind_password: SecretStr | None = None 

38 user_search_base: str = "" 

39 user_search_filter: str = "(sAMAccountName={username})" 

40 user_dn_attribute: str = "distinguishedName" 

41 group_search_base: str | None = None 

42 group_search_filter: str | None = None 

43 require_group_membership: str | None = None 

44 tls_ca_cert_file: str | None = None 

45 tls_cert_file: str | None = None 

46 tls_key_file: str | None = None 

47 timeout: int = 30 

48 max_connections: int = 10 

49 

50 

51class LDAPManager: 

52 """Manager for LDAP/Active Directory authentication operations. 

53 

54 Handles user authentication, user lookups, and group membership validation 

55 against LDAP/Active Directory servers. 

56 """ 

57 

58 def __init__( 

59 self, 

60 providers: dict[str, Any], # dict[str, LDAPProviderConfig] from admin 

61 http_client: HTTPClientProtocol | None = None, 

62 ): 

63 """Initialize LDAP manager. 

64 

65 Args: 

66 providers: Dictionary of LDAP provider configurations 

67 http_client: Optional HTTP client (not used for LDAP, but for consistency) 

68 """ 

69 if not LDAP_AVAILABLE: 

70 raise ImportError( 

71 "LDAP support requires 'ldap3' package. Install with: pip install ldap3", 

72 ) 

73 

74 self.providers = {} 

75 for name, config in providers.items(): 

76 # Convert from admin config to auth config 

77 self.providers[name] = LDAPProvider( 

78 name=config.name, 

79 server_url=config.server_url, 

80 bind_dn=config.bind_dn, 

81 bind_password=config.bind_password, 

82 user_search_base=config.user_search_base, 

83 user_search_filter=config.user_search_filter, 

84 user_dn_attribute=config.user_dn_attribute, 

85 group_search_base=config.group_search_base, 

86 group_search_filter=config.group_search_filter, 

87 require_group_membership=config.require_group_membership, 

88 tls_ca_cert_file=config.tls_ca_cert_file, 

89 tls_cert_file=config.tls_cert_file, 

90 tls_key_file=config.tls_key_file, 

91 timeout=config.timeout, 

92 max_connections=config.max_connections, 

93 ) 

94 

95 self.http_client = http_client 

96 self._connection_pool: dict[str, list[Connection]] = {} 

97 

98 def __repr__(self) -> str: 

99 """Return developer-friendly string representation.""" 

100 return f"LDAPManager(providers={list(self.providers)!r})" 

101 

102 async def authenticate_user( 

103 self, 

104 provider_name: str, 

105 username: str, 

106 password: str, 

107 ) -> dict[str, Any] | None: 

108 """Authenticate a user against LDAP. 

109 

110 Args: 

111 provider_name: Name of the LDAP provider 

112 username: Username to authenticate 

113 password: Password for authentication 

114 

115 Returns: 

116 User information dict if authentication successful, None otherwise 

117 

118 Raises: 

119 AuthError: If provider not found or LDAP error occurs 

120 """ 

121 provider = self.providers.get(provider_name) 

122 if not provider: 

123 raise AuthError(f"LDAP provider '{provider_name}' not found") 

124 

125 try: 

126 # Get user DN first 

127 user_dn = await self._get_user_dn(provider, username) 

128 if not user_dn: 

129 logger.warning("User %s not found in LDAP", username) 

130 return None 

131 

132 # Attempt authentication with user DN and password 

133 if await self._bind_with_credentials(provider, user_dn, password): 

134 # Get full user info 

135 user_info = await self._get_user_attributes(provider, user_dn) 

136 

137 # Check group membership if required 

138 if ( 

139 provider.require_group_membership 

140 and not await self.check_group_membership( 

141 provider_name, 

142 username, 

143 provider.require_group_membership, 

144 ) 

145 ): 

146 logger.warning( 

147 "User '%s' not member of required group '%s'", 

148 username, 

149 provider.require_group_membership, 

150 ) 

151 return None 

152 

153 return user_info 

154 logger.warning("LDAP authentication failed for user %s", username) 

155 return None 

156 

157 except LDAPException as e: 

158 logger.exception("LDAP error during authentication") 

159 raise AuthError(f"LDAP authentication failed: {e!s}") from e 

160 except (OSError, ConnectionError, ValueError) as e: 

161 logger.exception("Unexpected error during LDAP authentication") 

162 raise AuthError(f"LDAP authentication failed: {e!s}") from e 

163 

164 async def get_user_info( 

165 self, 

166 provider_name: str, 

167 username: str, 

168 ) -> dict[str, Any] | None: 

169 """Get user information from LDAP without authentication. 

170 

171 Args: 

172 provider_name: Name of the LDAP provider 

173 username: Username to look up 

174 

175 Returns: 

176 User information dict or None if not found 

177 """ 

178 provider = self.providers.get(provider_name) 

179 if not provider: 

180 raise AuthError(f"LDAP provider '{provider_name}' not found") 

181 

182 try: 

183 user_dn = await self._get_user_dn(provider, username) 

184 if not user_dn: 

185 return None 

186 

187 return await self._get_user_attributes(provider, user_dn) 

188 

189 except (LDAPException, ConnectionError, OSError): 

190 logger.exception("Error getting user info") 

191 return None 

192 

193 async def check_group_membership( 

194 self, 

195 provider_name: str, 

196 username: str, 

197 group_name: str, 

198 ) -> bool: 

199 """Check if user is a member of the specified group. 

200 

201 Args: 

202 provider_name: Name of the LDAP provider 

203 username: Username to check 

204 group_name: Group name to check membership for 

205 

206 Returns: 

207 True if user is a member of the group 

208 """ 

209 provider = self.providers.get(provider_name) 

210 if not provider: 

211 raise AuthError(f"LDAP provider '{provider_name}' not found") 

212 

213 try: 

214 # Get user DN 

215 user_dn = await self._get_user_dn(provider, username) 

216 if not user_dn: 

217 return False 

218 

219 # Get group DN 

220 group_dn = await self._get_group_dn(provider, group_name) 

221 if not group_dn: 

222 return False 

223 

224 # Check membership 

225 conn = await self._get_connection(provider) 

226 try: 

227 # Search for group and check member attribute 

228 conn.search( 

229 search_base=group_dn, 

230 search_filter="(objectClass=*)", 

231 search_scope=BASE, 

232 attributes=["member", "memberOf"], 

233 ) 

234 

235 if conn.entries: 

236 entry = conn.entries[0] 

237 members = entry.member.values if hasattr(entry, "member") else [] 

238 

239 # Check if user DN is in members 

240 if user_dn in members: 

241 return True 

242 

243 # Also check memberOf on user (for some LDAP implementations) 

244 conn.search( 

245 search_base=user_dn, 

246 search_filter="(objectClass=*)", 

247 search_scope=BASE, 

248 attributes=["memberOf"], 

249 ) 

250 

251 if conn.entries: 

252 user_entry = conn.entries[0] 

253 member_of = ( 

254 user_entry.memberOf.values 

255 if hasattr(user_entry, "memberOf") 

256 else [] 

257 ) 

258 return group_dn in member_of 

259 

260 return False 

261 

262 finally: 

263 self._return_connection(provider, conn) 

264 

265 except (LDAPException, ConnectionError, OSError): 

266 logger.exception("Error checking group membership") 

267 return False 

268 

269 async def _get_user_dn(self, provider: LDAPProvider, username: str) -> str | None: 

270 """Get the DN for a username by searching LDAP. 

271 

272 Args: 

273 provider: LDAP provider configuration 

274 username: Username to search for 

275 

276 Returns: 

277 User DN or None if not found 

278 """ 

279 conn = await self._get_connection(provider) 

280 try: 

281 search_filter = provider.user_search_filter.format(username=username) 

282 

283 conn.search( 

284 search_base=provider.user_search_base, 

285 search_filter=search_filter, 

286 search_scope=SUBTREE, 

287 attributes=[provider.user_dn_attribute], 

288 ) 

289 

290 if conn.entries: 

291 entry = conn.entries[0] 

292 value = getattr(entry, provider.user_dn_attribute).value 

293 # Normalize to str when possible 

294 if isinstance(value, bytes): 

295 return value.decode("utf-8", errors="ignore") 

296 if isinstance(value, str): 

297 return value 

298 if value is not None: 

299 return str(value) 

300 

301 return None 

302 

303 finally: 

304 self._return_connection(provider, conn) 

305 

306 async def _get_group_dn( 

307 self, 

308 provider: LDAPProvider, 

309 group_name: str, 

310 ) -> str | None: 

311 """Get the DN for a group by searching LDAP. 

312 

313 Args: 

314 provider: LDAP provider configuration 

315 group_name: Group name to search for 

316 

317 Returns: 

318 Group DN or None if not found 

319 """ 

320 if not provider.group_search_base or not provider.group_search_filter: 

321 return None 

322 

323 conn = await self._get_connection(provider) 

324 try: 

325 search_filter = provider.group_search_filter.format(group=group_name) 

326 

327 conn.search( 

328 search_base=provider.group_search_base, 

329 search_filter=search_filter, 

330 search_scope=SUBTREE, 

331 attributes=["distinguishedName"], 

332 ) 

333 

334 if conn.entries: 

335 entry = conn.entries[0] 

336 value = entry.distinguishedName.value 

337 if isinstance(value, bytes): 

338 return value.decode("utf-8", errors="ignore") 

339 if isinstance(value, str): 

340 return value 

341 if value is not None: 

342 return str(value) 

343 

344 return None 

345 

346 finally: 

347 self._return_connection(provider, conn) 

348 

349 async def _bind_with_credentials( 

350 self, 

351 provider: LDAPProvider, 

352 user_dn: str, 

353 password: str, 

354 ) -> bool: 

355 """Attempt to bind to LDAP with user credentials. 

356 

357 Args: 

358 provider: LDAP provider configuration 

359 user_dn: User DN to bind with 

360 password: Password for binding 

361 

362 Returns: 

363 True if bind successful 

364 """ 

365 # Create a new connection for authentication 

366 server = Server(provider.server_url, get_info=ALL) 

367 

368 try: 

369 conn = Connection( 

370 server, 

371 user=user_dn, 

372 password=password, 

373 auto_bind=True, 

374 read_only=True, 

375 receive_timeout=provider.timeout, 

376 ) 

377 

378 # Bind was successful if we get here 

379 conn.unbind() 

380 return True 

381 

382 except LDAPException: 

383 return False 

384 

385 async def _get_user_attributes( 

386 self, 

387 provider: LDAPProvider, 

388 user_dn: str, 

389 ) -> dict[str, Any]: 

390 """Get user attributes from LDAP. 

391 

392 Args: 

393 provider: LDAP provider configuration 

394 user_dn: User DN to get attributes for 

395 

396 Returns: 

397 Dictionary of user attributes 

398 """ 

399 conn = await self._get_connection(provider) 

400 try: 

401 # Common attributes to retrieve 

402 attributes = [ 

403 "objectGUID", 

404 "objectSid", 

405 "sAMAccountName", 

406 "userPrincipalName", 

407 "mail", 

408 "displayName", 

409 "givenName", 

410 "sn", 

411 "cn", 

412 "distinguishedName", 

413 "memberOf", 

414 "userAccountControl", 

415 "whenCreated", 

416 "whenChanged", 

417 ] 

418 

419 conn.search( 

420 search_base=user_dn, 

421 search_filter="(objectClass=*)", 

422 search_scope=BASE, 

423 attributes=attributes, 

424 ) 

425 

426 if conn.entries: 

427 entry = conn.entries[0] 

428 user_info = {} 

429 

430 for attr in attributes: 

431 if hasattr(entry, attr): 

432 value = getattr(entry, attr).value 

433 # Convert bytes to string for certain attributes 

434 if isinstance(value, bytes): 

435 if attr in ["objectGUID", "objectSid"]: 

436 # Keep as bytes for unique identifiers 

437 user_info[attr] = value.hex() 

438 else: 

439 user_info[attr] = value.decode("utf-8", errors="ignore") 

440 else: 

441 user_info[attr] = value 

442 

443 return user_info 

444 

445 return {} 

446 

447 finally: 

448 self._return_connection(provider, conn) 

449 

450 async def _get_connection(self, provider: LDAPProvider) -> Connection: 

451 """Get a connection from the pool or create a new one. 

452 

453 Args: 

454 provider: LDAP provider configuration 

455 

456 Returns: 

457 LDAP connection 

458 """ 

459 provider_name = provider.name 

460 

461 # Initialize pool if needed 

462 if provider_name not in self._connection_pool: 

463 self._connection_pool[provider_name] = [] 

464 

465 # Try to get existing connection 

466 if self._connection_pool[provider_name]: 

467 conn = self._connection_pool[provider_name].pop() 

468 if conn.bound: 

469 return conn 

470 

471 # Create new connection 

472 server = Server(provider.server_url, get_info=ALL) 

473 

474 return Connection( 

475 server, 

476 user=provider.bind_dn, 

477 password=( 

478 provider.bind_password.get_secret_value() 

479 if provider.bind_password is not None 

480 else None 

481 ), 

482 auto_bind=True, 

483 read_only=True, 

484 receive_timeout=provider.timeout, 

485 ) 

486 

487 def _return_connection(self, provider: LDAPProvider, conn: Connection) -> None: 

488 """Return a connection to the pool. 

489 

490 Args: 

491 provider: LDAP provider configuration 

492 conn: Connection to return 

493 """ 

494 provider_name = provider.name 

495 

496 if provider_name not in self._connection_pool: 

497 self._connection_pool[provider_name] = [] 

498 

499 # Only return if pool not full and connection is still valid 

500 if ( 

501 len(self._connection_pool[provider_name]) < provider.max_connections 

502 and conn.bound 

503 ): 

504 self._connection_pool[provider_name].append(conn) 

505 else: 

506 conn.unbind() 

507 

508 

509__all__ = ["LDAPManager", "LDAPProvider"]