Coverage for src/lexigram/auth/storage/oauth_identity_store/_mongodb.py: 26%

81 statements  

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

1"""MongoDB-backed OAuth identity store.""" 

2 

3from __future__ import annotations 

4 

5from datetime import datetime 

6from typing import TYPE_CHECKING, Any 

7 

8if TYPE_CHECKING: 

9 from lexigram.contracts.data import DatabaseProviderProtocol 

10 

11 

12from lexigram.di.decorators import inject 

13from lexigram.logging import get_logger 

14 

15logger = get_logger(__name__) 

16 

17from lexigram.auth.storage.oauth_identity_store._protocol import ( 

18 OAuthIdentity, 

19 OAuthIdentityStore, 

20) 

21 

22 

23@inject 

24class MongoDBOAuthIdentityStore(OAuthIdentityStore): 

25 """MongoDB-backed OAuth identity store""" 

26 

27 def __init__( 

28 self, 

29 db_provider: DatabaseProviderProtocol, 

30 collection_name: str = "oauth_identities", 

31 ): 

32 self.db_provider = db_provider 

33 self.collection_name = collection_name 

34 self._initialized = False 

35 

36 async def _ensure_collection(self) -> None: 

37 """Ensure collection exists with indexes""" 

38 if self._initialized: 

39 return 

40 

41 collection = self.db_provider.db[self.collection_name] # type: ignore[attr-defined] 

42 

43 # Create indexes 

44 await collection.create_index([("user_id", 1)]) 

45 await collection.create_index( 

46 [("provider", 1), ("provider_user_id", 1)], 

47 unique=True, 

48 ) 

49 

50 self._initialized = True 

51 

52 async def _identity_from_doc(self, doc: dict[str, Any]) -> OAuthIdentity: 

53 """Convert MongoDB document to OAuthIdentity object""" 

54 return OAuthIdentity( 

55 user_id=doc["user_id"], 

56 provider=doc["provider"], 

57 provider_user_id=doc["provider_user_id"], 

58 created_at=doc.get("created_at"), 

59 updated_at=doc.get("updated_at"), 

60 ) 

61 

62 async def _doc_from_identity(self, identity: OAuthIdentity) -> dict[str, Any]: 

63 """Convert OAuthIdentity object to MongoDB document""" 

64 return { 

65 "user_id": identity.user_id, 

66 "provider": identity.provider, 

67 "provider_user_id": identity.provider_user_id, 

68 "created_at": identity.created_at, 

69 "updated_at": identity.updated_at, 

70 } 

71 

72 async def create_oauth_identity( 

73 self, 

74 user_id: str, 

75 provider: str, 

76 provider_user_id: str, 

77 ) -> OAuthIdentity: 

78 """Create OAuth identity link""" 

79 await self._ensure_collection() 

80 

81 identity = OAuthIdentity( 

82 user_id=user_id, 

83 provider=provider, 

84 provider_user_id=provider_user_id, 

85 created_at=datetime.now(), 

86 updated_at=datetime.now(), 

87 ) 

88 

89 doc = await self._doc_from_identity(identity) 

90 collection = self.db_provider.db[self.collection_name] # type: ignore[attr-defined] 

91 await collection.insert_one(doc) 

92 

93 logger.info("Created OAuth identity: user=%s, provider=%s", user_id, provider) 

94 return identity 

95 

96 async def get_oauth_identity( 

97 self, 

98 provider: str, 

99 provider_user_id: str, 

100 ) -> OAuthIdentity | None: 

101 """Get OAuth identity by provider and provider user ID""" 

102 await self._ensure_collection() 

103 

104 collection = self.db_provider.db[self.collection_name] # type: ignore[attr-defined] 

105 doc = await collection.find_one( 

106 {"provider": provider, "provider_user_id": provider_user_id}, 

107 ) 

108 

109 return await self._identity_from_doc(doc) if doc else None 

110 

111 async def get_oauth_identities_for_user( 

112 self, 

113 user_id: str, 

114 ) -> list[OAuthIdentity]: 

115 """Get all OAuth identities for a user""" 

116 await self._ensure_collection() 

117 

118 collection = self.db_provider.db[self.collection_name] # type: ignore[attr-defined] 

119 cursor = collection.find({"user_id": user_id}) 

120 

121 identities = [] 

122 async for doc in cursor: 

123 identity = await self._identity_from_doc(doc) 

124 identities.append(identity) 

125 

126 return identities 

127 

128 async def delete_oauth_identity( 

129 self, 

130 provider: str, 

131 provider_user_id: str, 

132 ) -> bool: 

133 """Delete OAuth identity""" 

134 await self._ensure_collection() 

135 

136 collection = self.db_provider.db[self.collection_name] # type: ignore[attr-defined] 

137 result = await collection.delete_one( 

138 {"provider": provider, "provider_user_id": provider_user_id}, 

139 ) 

140 

141 deleted = bool(result.deleted_count > 0) 

142 if deleted: 

143 logger.info( 

144 "Deleted OAuth identity: provider=%s, user_id=%s", 

145 provider, 

146 provider_user_id, 

147 ) 

148 

149 return deleted 

150 

151 async def delete_oauth_identities_for_user(self, user_id: str) -> int: 

152 """Delete all OAuth identities for a user""" 

153 await self._ensure_collection() 

154 

155 collection = self.db_provider.db[self.collection_name] # type: ignore[attr-defined] 

156 result = await collection.delete_many({"user_id": user_id}) 

157 

158 deleted_count = int(result.deleted_count) 

159 if deleted_count > 0: 

160 logger.info( 

161 "Deleted %d OAuth identities for user %s", 

162 deleted_count, 

163 user_id, 

164 ) 

165 

166 return deleted_count 

167 

168 async def get_user_by_oauth_identity( 

169 self, 

170 provider: str, 

171 provider_user_id: str, 

172 ) -> str | None: 

173 """Get local user_id by OAuth provider and external user ID.""" 

174 await self._ensure_collection() 

175 

176 collection = self.db_provider.db[self.collection_name] # type: ignore[attr-defined] 

177 doc = await collection.find_one( 

178 {"provider": provider, "provider_user_id": provider_user_id}, 

179 {"user_id": 1}, 

180 ) 

181 

182 return doc.get("user_id") if doc else None 

183 

184 async def resolve_user_id( 

185 self, 

186 user_id_or_oauth_id: str, 

187 provider: str = "google", 

188 ) -> str | None: 

189 """Resolve user_id from either UUID or OAuth external ID. 

190 

191 Resolution logic: 

192 1. If user_id_or_oauth_id is a valid UUID format, check if user exists 

193 2. If not a valid UUID, treat it as an OAuth provider_user_id and look up 

194 """ 

195 import re 

196 

197 # Check if it's a valid UUID format 

198 uuid_pattern = re.compile( 

199 r"^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$", 

200 re.IGNORECASE, 

201 ) 

202 

203 is_uuid = bool(uuid_pattern.match(user_id_or_oauth_id)) 

204 

205 if is_uuid: 

206 # It's a UUID - check if user exists in users collection 

207 await self._ensure_collection() 

208 self.db_provider.db[self.collection_name] # type: ignore[attr-defined] 

209 

210 # We need to check the users collection - assume it's named "users" 

211 users_collection = self.db_provider.db["users"] # type: ignore[attr-defined] 

212 user_doc = await users_collection.find_one( 

213 {"_id": user_id_or_oauth_id}, 

214 {"_id": 1}, 

215 ) 

216 

217 return user_doc.get("_id") if user_doc else None 

218 # Not a UUID - treat as OAuth external ID 

219 return await self.get_user_by_oauth_identity(provider, user_id_or_oauth_id) 

220 

221 def resolve_user_id_sync( 

222 self, 

223 external_id: str, 

224 provider: str = "google", 

225 ) -> str | None: 

226 """Synchronous resolution is not supported for database-backed store.""" 

227 logger.warning( 

228 "sync_resolution_not_supported", 

229 provider=provider, 

230 external_id=external_id, 

231 store="MongoDBOAuthIdentityStore", 

232 ) 

233 return None