graphxr-database-proxy 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (33) hide show
  1. graphxr_database_proxy/__init__.py +16 -0
  2. graphxr_database_proxy/api/__init__.py +1 -0
  3. graphxr_database_proxy/api/database.py +266 -0
  4. graphxr_database_proxy/api/google.py +437 -0
  5. graphxr_database_proxy/api/projects.py +99 -0
  6. graphxr_database_proxy/common/util.py +38 -0
  7. graphxr_database_proxy/drivers/__init__.py +1 -0
  8. graphxr_database_proxy/drivers/base.py +56 -0
  9. graphxr_database_proxy/drivers/factory.py +36 -0
  10. graphxr_database_proxy/drivers/spanner.py +814 -0
  11. graphxr_database_proxy/main.py +110 -0
  12. graphxr_database_proxy/models/__init__.py +1 -0
  13. graphxr_database_proxy/models/google.py +50 -0
  14. graphxr_database_proxy/models/project.py +170 -0
  15. graphxr_database_proxy/proxy.py +495 -0
  16. graphxr_database_proxy/proxyForDev.py +290 -0
  17. graphxr_database_proxy/services/__init__.py +1 -0
  18. graphxr_database_proxy/services/project_service.py +150 -0
  19. graphxr_database_proxy/static/favicon.ico +0 -0
  20. graphxr_database_proxy/static/index.html +1 -0
  21. graphxr_database_proxy/static/main.7391ee9773c403483393.css +175 -0
  22. graphxr_database_proxy/static/main.7391ee9773c403483393.css.map +1 -0
  23. graphxr_database_proxy/static/main.ce3fbb85a7bc9452edb9.js +2 -0
  24. graphxr_database_proxy/static/main.ce3fbb85a7bc9452edb9.js.map +1 -0
  25. graphxr_database_proxy/static/vendors.70542a99f336a8021013.js +3 -0
  26. graphxr_database_proxy/static/vendors.70542a99f336a8021013.js.LICENSE.txt +95 -0
  27. graphxr_database_proxy/static/vendors.70542a99f336a8021013.js.map +1 -0
  28. graphxr_database_proxy-1.0.0.dist-info/METADATA +180 -0
  29. graphxr_database_proxy-1.0.0.dist-info/RECORD +33 -0
  30. graphxr_database_proxy-1.0.0.dist-info/WHEEL +5 -0
  31. graphxr_database_proxy-1.0.0.dist-info/entry_points.txt +2 -0
  32. graphxr_database_proxy-1.0.0.dist-info/licenses/LICENSE +21 -0
  33. graphxr_database_proxy-1.0.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,437 @@
1
+ # -*- coding: utf-8 -*-
2
+ """
3
+ Google Cloud API endpoints for listing projects and instances
4
+ """
5
+
6
+ import os
7
+
8
+ # Import and enable proxy interceptor
9
+ try:
10
+ from .. import proxyForDev
11
+ print("[OK] Proxy interceptor enabled")
12
+ except ImportError:
13
+ print("[INFO] Proxy interceptor not found, continuing without it")
14
+ pass
15
+
16
+ # Set environment variables to reduce Google Cloud client warning logs
17
+ os.environ.setdefault('GOOGLE_CLOUD_DISABLE_GRPC_FOR_GAE', 'true')
18
+ os.environ.setdefault('GRPC_VERBOSITY', 'ERROR')
19
+
20
+ from typing import Any, Dict, List
21
+ from fastapi import APIRouter, HTTPException, Request
22
+ from fastapi.responses import HTMLResponse, RedirectResponse
23
+ from google.cloud import spanner
24
+ from google.cloud import resourcemanager_v3
25
+ from google.oauth2 import service_account
26
+ import google.oauth2.credentials
27
+
28
+ import requests
29
+ from ..models.google import GoogleProject, SpannerDatabase
30
+ from google.api_core.exceptions import GoogleAPIError
31
+ from ..common.util import get_default_oauth_config
32
+
33
+ router = APIRouter(tags=["google"])
34
+
35
+
36
+ def get_spanner_client(project_id, credentials):
37
+ """Get Spanner client for given project and credentials"""
38
+ try:
39
+ spanner_client = spanner.Client(project=project_id, credentials=credentials)
40
+ return spanner_client
41
+ except Exception as e:
42
+ raise HTTPException(status_code=500, detail=f"Failed to create Spanner client: {e}")
43
+
44
+ def get_google_credentials(auth_info):
45
+ """
46
+ Get Google credentials based on auth_info
47
+ If contains 'token', use OAuth2, otherwise use service account
48
+ """
49
+ if auth_info.get("token"):
50
+ credentials = google.oauth2.credentials.Credentials(
51
+ token=auth_info.get("token"),
52
+ token_uri="https://oauth2.googleapis.com/token",
53
+ client_id=auth_info.get("client_id"),
54
+ client_secret=auth_info.get("client_secret"),
55
+ scopes=[
56
+ "https://www.googleapis.com/auth/cloudplatformprojects.readonly",
57
+ "https://www.googleapis.com/auth/spanner.admin",
58
+ "https://www.googleapis.com/auth/spanner.data",
59
+ "https://www.googleapis.com/auth/userinfo.profile",
60
+ "https://www.googleapis.com/auth/userinfo.email",
61
+ ]
62
+
63
+ )
64
+ return credentials, None
65
+ else:
66
+ # Service Account method
67
+ if auth_info:
68
+ credentials = service_account.Credentials.from_service_account_info(
69
+ auth_info
70
+ )
71
+ project = auth_info.get('project_id')
72
+ return credentials, project
73
+ else:
74
+ raise HTTPException(status_code=400, detail="Missing service account info")
75
+
76
+
77
+ @router.post("/api/google/spanner/list_projects", response_model=List[GoogleProject])
78
+ async def list_google_projects(request: Request):
79
+ """List Google Cloud projects"""
80
+ try:
81
+ # Get authentication info from request body json
82
+ body = await request.json()
83
+ auth = body.get('auth', {})
84
+
85
+ default_oauth = get_default_oauth_config()
86
+ newAuth = { **default_oauth, **auth }
87
+
88
+ credentials, project_id = get_google_credentials(newAuth)
89
+
90
+ is_service_account = credentials and isinstance(credentials, service_account.Credentials)
91
+ # Use Resource Manager API to list projects
92
+ projects = []
93
+
94
+ if is_service_account and project_id:
95
+ projects.append(GoogleProject(
96
+ name=project_id,
97
+ id=project_id,
98
+ instances=[]
99
+ ))
100
+ else:
101
+ client = resourcemanager_v3.ProjectsClient(credentials=credentials)
102
+ search_request = resourcemanager_v3.SearchProjectsRequest()
103
+ page_result = client.search_projects(request=search_request)
104
+ for project in page_result:
105
+ projects.append(GoogleProject(
106
+ name=project.display_name or project.name,
107
+ id=project.project_id,
108
+ instances=[]
109
+ ))
110
+
111
+ # Only keep projects that contain Spanner instances
112
+ spanner_projects = []
113
+
114
+ import concurrent.futures
115
+
116
+ def check_project_has_spanner(project):
117
+ """Check if a project has Spanner instances"""
118
+ try:
119
+ spanner_client = spanner.Client(project=project.id, credentials=credentials)
120
+ instances = list(spanner_client.list_instances())
121
+ if instances:
122
+ # Add the instances {id,name} to project.instances
123
+ project.instances = [{"id": inst.name.split('/')[-1], "name": inst.display_name or inst.instance_id} for inst in instances]
124
+ return project
125
+ return None
126
+ except (GoogleAPIError, Exception) as e:
127
+ print(f"Error checking project {project.id}: {e}")
128
+ # Skip projects that can't be accessed or have no Spanner
129
+ return None
130
+
131
+ # Use ThreadPoolExecutor to check projects in parallel
132
+ with concurrent.futures.ThreadPoolExecutor(max_workers=10) as executor:
133
+ # Submit all tasks
134
+ future_to_project = {
135
+ executor.submit(check_project_has_spanner, project): project
136
+ for project in projects
137
+ }
138
+
139
+ # Collect results as they complete
140
+ for future in concurrent.futures.as_completed(future_to_project):
141
+ result = future.result()
142
+ if result:
143
+ spanner_projects.append(result)
144
+
145
+ return spanner_projects
146
+
147
+ except Exception as e:
148
+ raise HTTPException(status_code=500, detail=str(e))
149
+
150
+
151
+ @router.post("/api/google/spanner/list_databases", response_model=list[SpannerDatabase])
152
+ async def list_google_databases(
153
+ request_data: Dict[str, Any],
154
+ ):
155
+ """List Google Cloud Spanner databases"""
156
+ try:
157
+ auth = request_data.get('auth', {})
158
+ project_id = auth.get('project_id')
159
+ instance_id = auth.get('instance_id')
160
+ default_oauth = get_default_oauth_config()
161
+ newAuth = { **default_oauth, **auth }
162
+ credentials, _ = get_google_credentials(newAuth)
163
+
164
+ if not project_id:
165
+ raise HTTPException(status_code=400, detail="Project ID not found")
166
+
167
+ # Use Spanner client
168
+ spanner_client = spanner.Client(project=project_id, credentials=credentials)
169
+
170
+ result = []
171
+
172
+ # List all Spanner instances
173
+ spanner_client = get_spanner_client(project_id, credentials)
174
+ instance = spanner_client.instance(instance_id)
175
+ if not instance.exists():
176
+ raise HTTPException(status_code=400, detail="Instance not found")
177
+
178
+ # List all databases in the instance
179
+ for database in instance.list_databases():
180
+ database_id = database.name.split('/')[-1]
181
+ database_name = database_id # Spanner databases usually don't have display names
182
+ databaseItem = {
183
+ "id": database_id,
184
+ "name": database_name,
185
+ "graphDBs": []
186
+ }
187
+
188
+ result.append(databaseItem)
189
+
190
+ # Check if there are Property Graph databases
191
+ try:
192
+ # Simple temporary query to get graph database information
193
+ try:
194
+ db = spanner_client.instance(instance_id).database(database_id)
195
+ # Use snapshot() or create session
196
+ with db.snapshot() as snapshot:
197
+ query = "SELECT PROPERTY_GRAPH_NAME FROM INFORMATION_SCHEMA.PROPERTY_GRAPHS"
198
+ results = snapshot.execute_sql(query)
199
+
200
+ for row in results:
201
+ graph_id = row[0]
202
+ databaseItem["graphDBs"].append({
203
+ "id": graph_id,
204
+ "name": graph_id
205
+ })
206
+ except:
207
+ # Skip if query fails
208
+ pass
209
+
210
+ except Exception as graph_error:
211
+ # If query fails, ignore graph database information
212
+ print(f"Could not query graph databases: {graph_error}")
213
+ pass
214
+
215
+ return result
216
+
217
+ except Exception as e:
218
+ raise HTTPException(status_code=500, detail=str(e))
219
+
220
+
221
+ @router.get("/google/spanner/callback")
222
+ async def google_spanner_callback(request: Request):
223
+ try:
224
+ host = request.headers.get('host')
225
+ print(f"Login request from host: {host}")
226
+
227
+ if(host.startswith("localhost") == False):
228
+ raise HTTPException(status_code=400, detail="OAuth login only allowed from localhost")
229
+
230
+ default_oauth = get_default_oauth_config()
231
+ client_id = default_oauth.get("client_id")
232
+ client_secret = default_oauth.get("client_secret")
233
+
234
+ # Get query parameters from the request
235
+ code = request.query_params.get("code")
236
+ state = request.query_params.get("state")
237
+ error = request.query_params.get("error")
238
+
239
+ if error:
240
+ raise HTTPException(status_code=400, detail=f"OAuth error: {error}")
241
+
242
+ if not code:
243
+ raise HTTPException(status_code=400, detail="Authorization code not found")
244
+
245
+ # Exchange authorization code for access token
246
+ redirect_uri = f"http://{host}/google/spanner/callback"
247
+ token_url = "https://oauth2.googleapis.com/token"
248
+
249
+ token_data = {
250
+ "client_id": client_id,
251
+ "client_secret": client_secret,
252
+ "code": code,
253
+ "grant_type": "authorization_code",
254
+ "redirect_uri": redirect_uri
255
+ }
256
+
257
+ response = requests.post(token_url, data=token_data)
258
+
259
+ if response.status_code != 200:
260
+ raise HTTPException(status_code=500, detail=f"Token exchange failed: {response.text}")
261
+
262
+ token_info = response.json()
263
+ token = token_info.get("access_token")
264
+ id_token = token_info.get("id_token")
265
+ refresh_token = token_info.get("refresh_token") # Add refresh_token
266
+ expires_in = token_info.get("expires_in")
267
+
268
+ if not token:
269
+ raise HTTPException(status_code=500, detail="Access token not found in token response")
270
+
271
+ # get user info
272
+ userinfo_response = requests.get(
273
+ "https://www.googleapis.com/oauth2/v3/userinfo",
274
+ headers={"Authorization": f"Bearer {token}"}
275
+ )
276
+
277
+ if userinfo_response.status_code != 200:
278
+ raise HTTPException(status_code=500, detail=f"Userinfo request failed: {userinfo_response.text}")
279
+
280
+ user_info = userinfo_response.json()
281
+ email = user_info.get("email", "unknown")
282
+ return HTMLResponse(content=f'''
283
+ <html>
284
+ <head>
285
+ <title>Google OAuth2 Callback</title>
286
+ </head>
287
+ <body style="font-family: Arial, sans-serif; text-align: center; margin-top: 50px;">
288
+ <h1>Google OAuth2 Successful</h1>
289
+ <p>You can close this window.</p>
290
+ <p>Current Login Email: {email}</p>
291
+ {f'<p style="color: orange;">[WARN] Note: refresh_token not obtained, need to re-login after token expires</p>' if not refresh_token else '<p style="color: green;">[OK] refresh_token obtained, supports automatic refresh</p>'}
292
+ <script>
293
+ // Use new token management system
294
+ const tokenData = {{
295
+ "access_token": "{token}",
296
+ "refresh_token": "{refresh_token or ''}",
297
+ "id_token": "{id_token or ''}",
298
+ "expires_in": {expires_in or 3600},
299
+ "email": "{email}"
300
+ }};
301
+
302
+ // Save to new format
303
+ localStorage.setItem('g_auth_info', JSON.stringify({{
304
+ ...tokenData,
305
+ "expires_at": Date.now() + (tokenData.expires_in * 1000)
306
+ }}));
307
+
308
+ // Compatible with old storage method
309
+ localStorage.setItem("g_auth_token", "{token}");
310
+ localStorage.setItem("g_auth_refresh_token", "{refresh_token or ''}");
311
+ localStorage.setItem("g_auth_id_token", "{id_token or ''}");
312
+ localStorage.setItem("g_auth_expires_in", "{expires_in}");
313
+ localStorage.setItem("g_auth_state", "{state}");
314
+ localStorage.setItem("g_auth_email", "{email}");
315
+
316
+ // Display token validity information
317
+ console.log('[OK] Authentication successful');
318
+ console.log('[INFO] Token validity:', tokenData.expires_in, 'seconds');
319
+ console.log('[INFO] Refresh token:', tokenData.refresh_token ? 'available' : 'not available');
320
+ console.log('[INFO] Expiry time:', new Date(Date.now() + tokenData.expires_in * 1000).toLocaleString());
321
+
322
+ if (!tokenData.refresh_token) {{
323
+ console.warn('[WARN] Warning: refresh_token not obtained, need to re-login after token expires');
324
+ }}
325
+
326
+ setTimeout(() => {{
327
+ window.close();
328
+ }}, 3000); // Extended display time for user to see warning information
329
+ </script>
330
+ </body>
331
+ </html>
332
+ ''')
333
+ except Exception as e:
334
+ return HTMLResponse(content=f'''
335
+ <html>
336
+ <body>
337
+ <h1>Google OAuth2 Error</h1>
338
+ <p>{str(e)}</p>
339
+ <script>
340
+ console.error("OAuth callback error: {str(e)}");
341
+ window.close();
342
+ </script>
343
+ </body>
344
+ </html>
345
+ ''')
346
+
347
+
348
+ @router.get("/google/spanner/login")
349
+ async def google_spanner_login(request: Request):
350
+ """Initiate Google OAuth2 login flow for Spanner access"""
351
+ try:
352
+ host = request.headers.get('host')
353
+ if(host.startswith("localhost") == False):
354
+ raise HTTPException(status_code=400, detail="OAuth login only allowed from localhost")
355
+ # Read from config/default.google.localhost.auth.json
356
+ try:
357
+ default_oauth = get_default_oauth_config()
358
+ client_id = default_oauth.get("client_id")
359
+
360
+ if not client_id:
361
+ raise HTTPException(status_code=500, detail="client_id not found in config file")
362
+ except Exception as e:
363
+ raise HTTPException(status_code=500, detail=f"OAuth configuration error: {e}")
364
+
365
+ scopes = [
366
+ "https://www.googleapis.com/auth/cloudplatformprojects.readonly",
367
+ "https://www.googleapis.com/auth/spanner.admin",
368
+ "https://www.googleapis.com/auth/spanner.data",
369
+ "https://www.googleapis.com/auth/userinfo.profile",
370
+ "https://www.googleapis.com/auth/userinfo.email",
371
+ ]
372
+ redirect_uri = f"http://{host}/google/spanner/callback"
373
+
374
+ # auto redirect to the OAuth2 authorization URL
375
+ auth_url = (
376
+ f"https://accounts.google.com/o/oauth2/v2/auth?"
377
+ f"client_id={client_id}&"
378
+ f"response_type=code&"
379
+ f"scope={'+'.join(scopes)}&"
380
+ f"access_type=offline&"
381
+ f"prompt=consent&"
382
+ f"redirect_uri={redirect_uri}&"
383
+ f"state=spanner"
384
+ )
385
+
386
+ return RedirectResponse(url=auth_url)
387
+
388
+ except Exception as e:
389
+ return {"error": str(e)}
390
+
391
+
392
+ @router.post("/api/google/refresh-token")
393
+ async def refresh_google_token(request: Request):
394
+ """Refresh Google OAuth2 access token"""
395
+ try:
396
+ body = await request.json()
397
+ refresh_token = body.get("refresh_token")
398
+
399
+ if not refresh_token:
400
+ raise HTTPException(status_code=400, detail="Refresh token is required")
401
+
402
+ # Get OAuth configuration
403
+ default_oauth = get_default_oauth_config()
404
+ client_id = default_oauth.get("client_id")
405
+ client_secret = default_oauth.get("client_secret")
406
+
407
+ if not client_id or not client_secret:
408
+ raise HTTPException(status_code=500, detail="OAuth configuration not found")
409
+
410
+ # Call Google's token refresh endpoint
411
+ token_url = "https://oauth2.googleapis.com/token"
412
+ token_data = {
413
+ "client_id": client_id,
414
+ "client_secret": client_secret,
415
+ "refresh_token": refresh_token,
416
+ "grant_type": "refresh_token"
417
+ }
418
+
419
+ response = requests.post(token_url, data=token_data)
420
+
421
+ if response.status_code != 200:
422
+ raise HTTPException(status_code=500, detail=f"Token refresh failed: {response.text}")
423
+
424
+ token_info = response.json()
425
+
426
+ # Return new access token information
427
+ return {
428
+ "access_token": token_info.get("access_token"),
429
+ "id_token": token_info.get("id_token"),
430
+ "expires_in": token_info.get("expires_in", 3600),
431
+ "token_type": token_info.get("token_type", "Bearer")
432
+ }
433
+
434
+ except HTTPException:
435
+ raise
436
+ except Exception as e:
437
+ raise HTTPException(status_code=500, detail=f"Token refresh error: {str(e)}")
@@ -0,0 +1,99 @@
1
+ """
2
+ Project management API endpoints
3
+ """
4
+
5
+ from typing import List
6
+ from fastapi import APIRouter, HTTPException, Depends
7
+ from ..models.project import Project, ProjectCreate, ProjectUpdate
8
+ from ..services.project_service import ProjectService
9
+
10
+ router = APIRouter(prefix="/api/project", tags=["projects"])
11
+
12
+ # Dependency to get project service
13
+ def get_project_service() -> ProjectService:
14
+ return ProjectService()
15
+
16
+
17
+ @router.post("/create", response_model=Project)
18
+ async def create_project(
19
+ project_data: ProjectCreate,
20
+ service: ProjectService = Depends(get_project_service)
21
+ ):
22
+ """Create a new project"""
23
+ try:
24
+ # Check if project with same name already exists
25
+ existing = await service.get_project_by_name(project_data.name)
26
+ if existing:
27
+ raise HTTPException(
28
+ status_code=400,
29
+ detail=f"Project with name '{project_data.name}' already exists"
30
+ )
31
+
32
+ project = await service.create_project(project_data)
33
+ return project
34
+ except Exception as e:
35
+ raise HTTPException(status_code=500, detail=str(e))
36
+
37
+
38
+ @router.get("/list", response_model=List[Project])
39
+ async def list_projects(
40
+ service: ProjectService = Depends(get_project_service)
41
+ ):
42
+ """List all projects"""
43
+ try:
44
+ projects = await service.list_projects()
45
+ return projects
46
+ except Exception as e:
47
+ raise HTTPException(status_code=500, detail=str(e))
48
+
49
+
50
+ @router.get("/{project_id}", response_model=Project)
51
+ async def get_project(
52
+ project_id: str,
53
+ service: ProjectService = Depends(get_project_service)
54
+ ):
55
+ """Get a project by ID"""
56
+ try:
57
+ project = await service.get_project(project_id)
58
+ if not project:
59
+ raise HTTPException(status_code=404, detail="Project not found")
60
+ return project
61
+ except HTTPException:
62
+ raise
63
+ except Exception as e:
64
+ raise HTTPException(status_code=500, detail=str(e))
65
+
66
+
67
+ @router.put("/update", response_model=Project)
68
+ async def update_project(
69
+ project_id: str,
70
+ update_data: ProjectUpdate,
71
+ service: ProjectService = Depends(get_project_service)
72
+ ):
73
+ """Update a project"""
74
+ try:
75
+ project = await service.update_project(project_id, update_data)
76
+ if not project:
77
+ raise HTTPException(status_code=404, detail="Project not found")
78
+ return project
79
+ except HTTPException:
80
+ raise
81
+ except Exception as e:
82
+ raise HTTPException(status_code=500, detail=str(e))
83
+
84
+
85
+ @router.delete("/delete")
86
+ async def delete_project(
87
+ project_id: str,
88
+ service: ProjectService = Depends(get_project_service)
89
+ ):
90
+ """Delete a project"""
91
+ try:
92
+ success = await service.delete_project(project_id)
93
+ if not success:
94
+ raise HTTPException(status_code=404, detail="Project not found")
95
+ return {"message": "Project deleted successfully"}
96
+ except HTTPException:
97
+ raise
98
+ except Exception as e:
99
+ raise HTTPException(status_code=500, detail=str(e))
@@ -0,0 +1,38 @@
1
+ import os
2
+ import sys
3
+ import json
4
+ from fastapi import HTTPException
5
+
6
+ # Add project root directory to Python path for importing proxy modules
7
+ current_dir = os.path.dirname(os.path.abspath(__file__))
8
+ project_root = os.path.join(current_dir, '..', '..', '..')
9
+ sys.path.insert(0, project_root)
10
+
11
+ def read_json_file(file_path):
12
+ """Read JSON file with error handling"""
13
+ try:
14
+ if not os.path.exists(file_path):
15
+ raise FileNotFoundError(f"Config file not found: {file_path}")
16
+
17
+ with open(file_path, 'r', encoding='utf-8') as file:
18
+ return json.load(file)
19
+ except json.JSONDecodeError as e:
20
+ raise HTTPException(status_code=500, detail=f"Invalid JSON in config file: {e}")
21
+ except Exception as e:
22
+ raise HTTPException(status_code=500, detail=f"Error reading config file: {e}")
23
+
24
+ def get_default_oauth_config():
25
+ """Get default OAuth config from config/default.google.localhost.auth.json"""
26
+ try:
27
+ # Get the project root directory (assuming this file is in src/graphxr_database_proxy/api/)
28
+ current_dir = os.path.dirname(os.path.abspath(__file__))
29
+ project_root = os.path.join(current_dir, '..', '..', '..')
30
+ config_path = os.path.join(project_root, 'config', 'default.google.localhost.auth.json')
31
+ default_oauth = read_json_file(config_path).get("web", {})
32
+ client_id = default_oauth.get("client_id")
33
+ client_secret = default_oauth.get("client_secret")
34
+ if not client_id or not client_secret:
35
+ raise HTTPException(status_code=500, detail="client_id or client_secret not found in config file")
36
+ return default_oauth
37
+ except Exception as e:
38
+ raise HTTPException(status_code=500, detail=f"OAuth configuration error: {e}")
@@ -0,0 +1 @@
1
+ # Drivers package
@@ -0,0 +1,56 @@
1
+ """
2
+ Base driver interface for database connections
3
+ """
4
+
5
+ from abc import ABC, abstractmethod
6
+ from typing import Any, Dict, List, Optional
7
+ from ..models.project import Project ,DatabaseConfig, QueryResponse, SchemaResponse, GraphSchemaResponse, SampleDataResponse
8
+
9
+
10
+ class BaseDatabaseDriver(ABC):
11
+ """Base class for database drivers"""
12
+
13
+ def __init__(self, project: Project):
14
+ self.project = project
15
+ self.config = project.database_config
16
+ self._connection = None
17
+
18
+ @abstractmethod
19
+ async def connect(self) -> None:
20
+ """Establish connection to the database"""
21
+ pass
22
+
23
+ @abstractmethod
24
+ async def disconnect(self) -> None:
25
+ """Close connection to the database"""
26
+ pass
27
+
28
+ @abstractmethod
29
+ async def test_connection(self) -> bool:
30
+ """Test if connection is working"""
31
+ pass
32
+
33
+ @abstractmethod
34
+ async def execute_query(self, query: str, parameters: Dict[str, Any] = None) -> QueryResponse:
35
+ """Execute a query"""
36
+ pass
37
+
38
+ @abstractmethod
39
+ async def get_schema(self) -> SchemaResponse:
40
+ """Get database schema"""
41
+ pass
42
+
43
+ @abstractmethod
44
+ async def get_graph_schema(self) -> GraphSchemaResponse:
45
+ """Get graph database schema"""
46
+ pass
47
+
48
+ @abstractmethod
49
+ async def get_sample_data(self) -> SampleDataResponse:
50
+ """Get sample data from database"""
51
+ pass
52
+
53
+ @abstractmethod
54
+ def get_api_info(self, project_name: str) -> Dict[str, Any]:
55
+ """Get API information for this database"""
56
+ pass
@@ -0,0 +1,36 @@
1
+ """
2
+ Driver factory for creating database drivers
3
+ """
4
+
5
+ from typing import Dict, Type
6
+ from .base import BaseDatabaseDriver
7
+ from .spanner import SpannerDriver
8
+ from ..models.project import Project, DatabaseConfig, DatabaseType
9
+
10
+
11
+ class DriverFactory:
12
+ """Factory for creating database drivers"""
13
+
14
+ _drivers: Dict[DatabaseType, Type[BaseDatabaseDriver]] = {
15
+ DatabaseType.SPANNER: SpannerDriver,
16
+ }
17
+
18
+ @classmethod
19
+ def create_driver(cls, project: Project) -> BaseDatabaseDriver:
20
+ """Create a driver instance for the given database type"""
21
+ config = project.database_config
22
+ driver_class = cls._drivers.get(config.type)
23
+ if not driver_class:
24
+ raise ValueError(f"Unsupported database type: {config.type}")
25
+
26
+ return driver_class(project)
27
+
28
+ @classmethod
29
+ def register_driver(cls, db_type: DatabaseType, driver_class: Type[BaseDatabaseDriver]) -> None:
30
+ """Register a new driver type"""
31
+ cls._drivers[db_type] = driver_class
32
+
33
+ @classmethod
34
+ def get_supported_types(cls) -> list[DatabaseType]:
35
+ """Get list of supported database types"""
36
+ return list(cls._drivers.keys())