import re
from collections.abc import Iterator
import pandas as pd
from sqlalchemy import create_engine, inspect, text
from sqlalchemy.engine import URL
from app.connectors.base import BaseConnector
SAFE_OBJECT=re.compile(r'^[A-Za-z0-9_.$#\[\]-]+$')
class SQLAlchemyConnector(BaseConnector):
    def __init__(self,db_type,host,port,database,username,password,options=None):
        self.db_type=db_type; self.options=options or {}; driver={'postgresql':'postgresql+psycopg','mysql':'mysql+pymysql','oracle':'oracle+oracledb','sqlserver':'mssql+pyodbc'}[db_type]; query={}
        if db_type=='sqlserver': query={'driver':self.options.get('driver','ODBC Driver 18 for SQL Server'),'TrustServerCertificate':self.options.get('trust_server_certificate','no'),'Encrypt':self.options.get('encrypt','yes')}
        self.engine=create_engine(URL.create(driver,username=username,password=password,host=host,port=port,database=database,query=query),pool_pre_ping=True,pool_recycle=300)
    def test(self):
        with self.engine.connect() as c: c.execute(text('SELECT 1'))
        return True
    def catalog(self,schema=None):
        i=inspect(self.engine); schemas=[schema] if schema else i.get_schema_names(); out=[]
        for s in schemas:
            if s.lower() in {'information_schema','pg_catalog','sys'}: continue
            for kind,names in [('table',i.get_table_names(schema=s)),('view',i.get_view_names(schema=s))]:
                for n in names:
                    try: cols=[{'name':c['name'],'type':str(c['type']),'nullable':c.get('nullable',True)} for c in i.get_columns(n,schema=s)]
                    except Exception: cols=[]
                    out.append({'schema':s,'name':n,'qualified_name':f'{s}.{n}' if s else n,'type':kind,'columns':cols})
        return out
    def extract(self,source_object,incremental_column=None,watermark_value=None,chunk_size=50000):
        if not SAFE_OBJECT.fullmatch(source_object): raise ValueError('source_object inválido')
        if incremental_column and not SAFE_OBJECT.fullmatch(incremental_column): raise ValueError('incremental_column inválida')
        sql=f'SELECT * FROM {source_object}'; params={}
        if incremental_column and watermark_value is not None: sql+=f' WHERE {incremental_column} > :watermark'; params['watermark']=watermark_value
        if incremental_column: sql+=f' ORDER BY {incremental_column}'
        yield from pd.read_sql_query(text(sql),self.engine,params=params,chunksize=chunk_size)

    def preview(self, qualified_name: str, limit: int = 50):
        from sqlalchemy import text
        safe = qualified_name.replace('"','').replace(';','')
        limit = max(1, min(int(limit), 5000))
        dialect = self.engine.dialect.name
        if dialect in {'mssql'}:
            sql = f"SELECT TOP {limit} * FROM {safe}"
        elif dialect == 'oracle':
            sql = f"SELECT * FROM {safe} FETCH FIRST {limit} ROWS ONLY"
        else:
            sql = f"SELECT * FROM {safe} LIMIT {limit}"
        with self.engine.connect() as conn:
            result = conn.execute(text(sql))
            columns = list(result.keys())
            rows = [[v.isoformat() if hasattr(v,'isoformat') else v for v in row] for row in result.fetchall()]
        return {'columns':columns,'rows':rows}
