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
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:58 +0800
1"""MongoDB-backed OAuth identity store."""
3from __future__ import annotations
5from datetime import datetime
6from typing import TYPE_CHECKING, Any
8if TYPE_CHECKING:
9 from lexigram.contracts.data import DatabaseProviderProtocol
12from lexigram.di.decorators import inject
13from lexigram.logging import get_logger
15logger = get_logger(__name__)
17from lexigram.auth.storage.oauth_identity_store._protocol import (
18 OAuthIdentity,
19 OAuthIdentityStore,
20)
23@inject
24class MongoDBOAuthIdentityStore(OAuthIdentityStore):
25 """MongoDB-backed OAuth identity store"""
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
36 async def _ensure_collection(self) -> None:
37 """Ensure collection exists with indexes"""
38 if self._initialized:
39 return
41 collection = self.db_provider.db[self.collection_name] # type: ignore[attr-defined]
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 )
50 self._initialized = True
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 )
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 }
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()
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 )
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)
93 logger.info("Created OAuth identity: user=%s, provider=%s", user_id, provider)
94 return identity
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()
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 )
109 return await self._identity_from_doc(doc) if doc else None
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()
118 collection = self.db_provider.db[self.collection_name] # type: ignore[attr-defined]
119 cursor = collection.find({"user_id": user_id})
121 identities = []
122 async for doc in cursor:
123 identity = await self._identity_from_doc(doc)
124 identities.append(identity)
126 return identities
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()
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 )
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 )
149 return deleted
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()
155 collection = self.db_provider.db[self.collection_name] # type: ignore[attr-defined]
156 result = await collection.delete_many({"user_id": user_id})
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 )
166 return deleted_count
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()
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 )
182 return doc.get("user_id") if doc else None
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.
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
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 )
203 is_uuid = bool(uuid_pattern.match(user_id_or_oauth_id))
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]
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 )
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)
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