-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdb_interface.py
More file actions
119 lines (98 loc) · 4.35 KB
/
Copy pathdb_interface.py
File metadata and controls
119 lines (98 loc) · 4.35 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
import os
from abc import ABC, abstractmethod
from typing import List, Dict, Any, Tuple
from sqlalchemy import create_engine, inspect, text
from sqlalchemy.engine import Engine, Result
# ==========================================
# 🔧 DATABASE CONFIGURATION (MySQL — localhost only)
# Fill in YOUR local MySQL credentials below
# ==========================================
DB_HOST = "localhost"
DB_PORT = "3306"
DB_USER = "root" # ← Your MySQL username
DB_PASSWORD = "" # ← Your MySQL password
DB_NAME = "your_database_name" # ← Your database name
# ==========================================
class DatabaseConnector(ABC):
"""Abstract base class for DB connectors."""
def __init__(self, db_url: str):
self.db_url = db_url
self.engine: Engine = create_engine(db_url, future=True)
@abstractmethod
def get_schema(self) -> Dict[str, List[Dict[str, Any]]]:
"""Return a mapping {table_name: [{name: col, type: type, primary_key: bool}, ...]}"""
@abstractmethod
def execute_query(self, query: str) -> List[Tuple]:
"""Execute a SELECT query and return up to 100 rows as tuples."""
@abstractmethod
def get_relationships(self) -> List[str]:
"""Return foreign-key relationships in the form ``table.column -> ref_table.ref_column``."""
def _inspect(self):
return inspect(self.engine)
class SQLAlchemyConnector(DatabaseConnector):
"""Concrete implementation using SQLAlchemy for MySQL."""
def get_schema(self) -> Dict[str, List[Dict[str, Any]]]:
inspector = self._inspect()
schema: Dict[str, List[Dict[str, Any]]] = {}
for table in inspector.get_table_names():
columns = []
for col in inspector.get_columns(table):
columns.append(
{
"name": col["name"],
"type": str(col["type"]),
"primary_key": bool(col.get("primary_key", False)),
}
)
schema[table] = columns
return schema
def execute_query(self, query: str) -> List[Tuple]:
with self.engine.connect() as conn:
result: Result = conn.execute(text(query))
return [tuple(row) for row in result.fetchmany(100)]
def get_relationships(self) -> List[str]:
inspector = self._inspect()
rels: List[str] = []
for table in inspector.get_table_names():
fks = inspector.get_foreign_keys(table)
for fk in fks:
for src, dst in zip(fk["constrained_columns"], fk["referred_columns"]):
rels.append(f"{table}.{src} -> {fk['referred_table']}.{dst}")
return rels
# Helper factory ------------------------------------------------------------
def get_connector() -> DatabaseConnector:
"""Create a MySQL connector using the config variables above.
If the MySQL connection fails, fall back to an in‑memory SQLite database for testing.
"""
if DB_PASSWORD:
db_url = f"mysql+pymysql://{DB_USER}:{DB_PASSWORD}@{DB_HOST}:{DB_PORT}/{DB_NAME}"
else:
db_url = f"mysql+pymysql://{DB_USER}@{DB_HOST}:{DB_PORT}/{DB_NAME}"
try:
# Attempt MySQL connection
connector = SQLAlchemyConnector(db_url)
# Force a test connection to verify availability
with connector.engine.connect() as conn:
pass
return connector
except Exception as e:
# Fallback to SQLite in‑memory for environments without MySQL
print(f"MySQL connection failed ({e}), falling back to SQLite in‑memory for testing.")
fallback_url = "sqlite:///:memory:"
return SQLAlchemyConnector(fallback_url)
def format_schema(schema: dict) -> str:
"""Return a newline-separated representation of a schema dictionary.
Example: "employees(id INTEGER, name TEXT)" per line.
"""
lines = []
for table, cols in schema.items():
col_text = ", ".join(f"{c['name']} {c['type']}" for c in cols)
lines.append(f"{table}({col_text})")
return "\n".join(lines)
def format_relationships(relations: list) -> str:
"""Return a newline-separated list of relationship strings.
If the list is empty, return a placeholder message.
"""
if not relations:
return "No relationships found."
return "\n".join(relations)