#$Id: sa_generator.py 715 2007-01-25 22:09:36Z sdobrev $
#store+print SA constructor-args in order to recreate the source

import sqlalchemy
class CV( sqlalchemy.ClauseVisitor):
    ops = {
        '=':'==',
    }
    def __init__(me):
        me.stack = []
    def visit_bindparam(me, c):
        t = repr(c.value)
        me.stack.append( t)
    def visit_column(me, c):
        if isinstance( c.table, sqlalchemy.Table):
            varname= table_varname( c.table)
        else:
            varname= c.table.name
        t = varname + '.c.' + c.name
        me.stack.append( t)
    def visit_binary(me, b):
        r = me.stack.pop()
        l = me.stack.pop()
        t = ' '.join( [l, me.ops.get( b.operator, b.operator), r ])
        me.stack.append( t)

do_columns = False
level = 0
def tstr(o):
    global level
    level+=1
    if isinstance(o,type):  r = o.__name__
    elif isinstance(o,str): r = repr(o)
    elif isinstance(o,dict):
        r = ('\n'+16*' ').join( ['{']+ ['%r: %r,' % kv for kv in o.iteritems() ] + ['}'] )
#    elif do_columns and isinstance(o, sql.sqlalchemy.Column):
#        r = o.table.name+'.c.'+o.name
    elif (do_columns and isinstance(o, sqlalchemy.sql.ColumnElement) or
           isinstance(o, sqlalchemy.sql._BinaryClause) ):    #ClauseElement
        cv = CV()
        o.accept_visitor(cv)
#        print o, cv.stack
        r = cv.stack.pop()
        assert not cv.stack
    else:
#        print o, type(o)
        try: r = o.tstr
        except AttributeError:
            if isinstance( o, sqlalchemy.Alias):
                #print id(o), id(o.selectable), type( o.original)
                #print getattr(o, 'tstr','o'), getattr(o.selectable, 'tstr','osel'), getattr(o.original, 'tstr','org'), repr(o.selectable)
                try: return tstr( o.selectable)#.tstr
                except AttributeError: pass
            r = str(o)
#        if isinstance(o, sqlalchemy.ForeignKey):    #XXX HACK
#            if o.use_alter and 'use_alter' not in r:
#                r[-2:-2] = 'use_alter= True, name= '+repr(o.name)+','
    level-=1
    return r

def tstr2(o):
#    print type(o)
    try: return o.tstr
    except AttributeError: return o.org__repr__()

def repr2tstr( klas):
    klas.org__repr__ = klas.__repr__
    klas.__repr__ = tstr2

class Tstr:
    nl = ''
    nl4args = None
    no_kargs = {}   #ignore these kargs of these values
    def __init__( me, nl=None, nl4args=None, no_kargs =None):
        if nl is None: nl = me.nl
        if nl4args is None: nl4args = me.nl4args
        if no_kargs is None: no_kargs = me.no_kargs

        if nl4args is None: nl4args = nl
        me.nl = nl
        me.nl4args = nl4args
        me.no_kargs = no_kargs

    def thestr( me, tself, name, args, kargs):
        nl = me.nl
        nl4args = me.nl4args
        t = tself and (tself + '.') or ''
        ks = kargs.keys()
        ks.sort()
        return ( t + name+ '( '+
                nl.join(
                    [ nl4args.join( [str(tstr(a))+', ' for a in args]) ] +
                    [ (level*'  '+'%s= %s, ') % (k,tstr(kargs[k]))
                        for k in ks #kargs.iteritems()
                        if kargs[k] is not me.no_kargs.get( k,'anyany')
                    ] + [')'] )
                )
class TstrSelf( Tstr):
    def __init__( me, name, args, kargs, tself ='', **kargs4setup):
        me.name = name
        me.tself = tself
        me.args = args
        me.kargs = kargs
        Tstr.__init__( me, **kargs4setup)
    def __str__( me):
        return me.thestr( me.tself, me.name, me.args, me.kargs)

class duper( Tstr):
    def dup( me, self, *args,**kargs):
        base = me.base
        #do not compare things containing columns!
        t = me.thestr( self and tstr2(self) or '', #XXX move to Tstr..
                        base.__name__, args, kargs)
        if self is None:
            r = base( *args, **kargs)
        else:
            r = base( self, *args, **kargs)
        if me.otherstr:
            r.tstr2 = t
            t = me.otherstr(r)
        r.tstr = t
        return r

    def __init__( me, base, otherstr =None, **kargs4setup):
        me.base = base
        me.otherstr = otherstr
        Tstr.__init__( me, **kargs4setup)

    def __call__( me, *args,**kargs):
        return me.dup( None, *args,**kargs)

class duper2( duper):
#    with_attrs = {} #add these attributes as kargs (if not of these values)
    def thestr( me, tself, name, args, kargs):
#        for k,v in me.with_attrs.iteritems():
#            a = getattr(
        return TstrSelf( name, args, kargs, tself=tself,
                            nl= me.nl, nl4args= me.nl4args, no_kargs= me.no_kargs)

#names = {}
def table_varname(t): return 'table_'+t.name
def punion_varname(u): return u.name #'punion_'+
def mapper_varname(m): return 'mapper_'+m.class_.__name__+(m.non_primary and '1' or '')

class Printer:
    def __init__( me, filename=''):
        me.out = ''
    def nl(me):
        me.out += '\n'

    def pklas( me, klas, getprops =None, **kargs_ignore):
        base = klas.__bases__[0].__name__
        name = klas.__name__
        if getprops:
            props = getprops(klas)
        else:
            props = klas.props
        ptr = 'data2' in props and 'data2' or 'name'
        me.out += '''\
class %(name)s( %(base)s):
    props = %(props)s
    data = property( lambda me: me.%(ptr)s)
''' % locals()

    def pklasi( me, Base, namespace, **kargs):
        me.Base = Base
        for k,klas in namespace.iteritems():
            if not isinstance( klas, type) or not issubclass( klas, Base): continue
            me.pklas( klas, **kargs)
        me.nl()

    def ptabli( me, meta):
        global do_columns
        do_columns=False
        alltbl = meta.tables.values() #[t1,t2,t3, t11,t12]
        alltbl.sort( key=lambda t:t.name)
        ind = '\n    '
        for t in alltbl:
            name = t.name
            varname = table_varname( t)
            t.tstr = varname
#            names[ t] = varname
            me.out += '%(varname)s = Table( %(name)r, meta,' % locals()
            me.out += ind + ind.join( [str(tstr(c))+',' for c in t.columns ]) + '\n)'
            me.nl()
        me.out += '''
meta.create_all()

'''
        do_columns=True

    def pklasi_tabli( me, meta, Base, namespace ):
        me.pklasi( Base, namespace)
        me.ptabli( meta)
    def punion( me, pu, mapper):
        pu_tstr = pu.tstr2
        if 'HACK4inhtype':
            items = pu_tstr.split("':") #n_items+1
            n =0
            for i in items[1:]:
                if '.select(' in i or 'join(' in i:
                    n+=1
            if not n:
                typ = ''
                pu_tstr+= ' #concrete'
            elif n == len(items)-1:
                pu_tstr = pu_tstr.replace( "'atype',", 'None,')
                pu_tstr+= ' #tableinh'
            else:
                #pu_tstr.replace( "'atype',", 'NotImplementedError,')
                pu_tstr+= ' XXX  NotImplementedError - mixed tableinh and concrete; use polymunion.py'
        me.out += punion_varname(pu)+ ' = ' + pu_tstr
        me.nl()
    def pmapi( me, namespace):
        maps = [ m for m in namespace.itervalues() if isinstance(m,sqlalchemy.orm.Mapper)]
        maps.sort( key=lambda m:m.class_.__name__)
        for m in maps:
            pu = m.select_table
            if isinstance( pu,sqlalchemy.sql.Alias):  #CompoundSelect
                me.punion( pu, m)

            varname = mapper_varname( m)
            t2 = m.tstr2
            me.out += varname + ' = '+ t2
            me.nl()
            for k,p in m.properties.iteritems():
                t = tstr(p)
                me.out += '%(varname)s.add_property( %(k)r, %(t)s )\n' % locals()
            me.nl()
        me.nl()

    head = '''
from sa_gentestbase import *

class AB( Test_AB0):
'''

    tail = '''

if __name__ == '__main__':
    setup()
    unittest.main()
'''

    def populate( me, namespace):
        Base = me.Base
        r = '''
#populate
'''
        s = [ (k,m) for k,m in namespace.iteritems() if isinstance( m, Base)]
        s.sort()
        names = {}
        for (k,m) in s:
            r += k +' = ' + m.__class__.__name__+'()\n'
            names[ id(m)] = k
        for (k,m) in s:
            for a in m.props:
                v = getattr( m, a, None)
                if v:
                    if isinstance( v, Base): v = names[id(v)]
                    else: v = repr(v)
                    r += k +'.'+a + ' = ' + v + '\n'

        A = namespace['A']
        B = namespace['B']
        r+= '''
session = create_session()
session.save(a)
session.save(b)
session.flush()

sa = str(a)
sb = str(b)
sbmulti = [ '''
        b_mul = 'sb' + (namespace.get('b1') and ', str(b1)' or '')
        r+= b_mul + ' ]'

        r+= ('''
samulti = [ sa'''
                    + (namespace.get('a1') and ', str(a1)' or '')
                    + (issubclass( B,A) and ', '+b_mul or '')
                + ' ]' )

        r+= '''
me.query( session, A,B, table_A,table_B, a.id, b.id, sa,sb, samulti, sbmulti )
'''

        me.out += r

def duper4polymorphic_union( polymorphic_union):
    return duper( polymorphic_union, otherstr=punion_varname)
####################### now redefine these
Column = duper( sqlalchemy.Column)
ForeignKey= duper( sqlalchemy.ForeignKey)

class Table( sqlalchemy.Table):
    org__repr__ = sqlalchemy.Table.__repr__
    __repr__ = tstr2
    dup = duper( sqlalchemy.Table.select)
    def select( me, *args, **kargs):
        r = me.dup.dup( me, *args, **kargs)
#        print id(r)
        return r

jdup = duper( sqlalchemy.Join.select)
def select4join( me, *args, **kargs):
    r = jdup.dup( me, *args, **kargs)
    return r
sqlalchemy.Join.select = select4join

def repr2alias(me):
#    print 'xxx', id(me), id(me.original), id(me.selectable), type(me.original), str(me.original), 'eoxx'
    try: return me.original.tstr    #selectable
    except AttributeError:
        return tstr2(me)
sqlalchemy.Alias.org__repr__ = sqlalchemy.Alias.__repr__
sqlalchemy.Alias.__repr__ = repr2alias
repr2tstr( sqlalchemy.Select)
repr2tstr( sqlalchemy.Join)

mapper= duper( sqlalchemy.mapper, otherstr=mapper_varname, nl='\n'+12*' ', nl4args='', no_kargs= dict( concrete=False) )
join  = duper( sqlalchemy.join)
relation = duper( sqlalchemy.relation, nl='\n'+12*' ', nl4args='', no_kargs= dict( remote_side=None) )  #lazy=True,
polymorphic_union= duper4polymorphic_union( sqlalchemy.polymorphic_union )

# vim:ts=4:sw=4:expandtab
