from sympy import *
from sympy.abc import q,x,y
init_printing()

"""g can onyl have positive powers of q (no shearing"""
def invert(g,o):
    g0 = g.subs(q,0)
    g0inv = g0.inv()
    m = eye(sqrt(len(g)))
    gnorm = expand(g*g0inv-m)
    series = zeros(sqrt(len(g)))
    for i in range(0,o+1):
        series = series + m
        m = m*(-gnorm)
    ginv = g0inv*series
    return expand(ginv)

"""same condition as before"""
def gauge(a,g,o):
    diffg = diff(g,q)
    ginv = invert(g,o)
    newa = expand(ginv*a*g+ginv*diffg)
    return(truncatea(newa,o-2))

"""here, assumed to have quadratic pole"""
def truncatea(a,o):
    size = sqrt(len(a))
    b = zeros(size)
    for i in range(0,size):
        for j in range(0,size):
            e = a[i,j]
            eseries = expand((expand(e*q**2) + O(q**(o+3))).removeO()*q**(-2))
            b[i,j] = eseries
    return b

def mirror_of_p1():
    a = -q**(-2)*Matrix([[0,2],[2,0]])+ q**(-1)*Matrix([[0,0],[0,1]])
    eigth= Rational(1,8)
    g = Matrix([[1,1],[1,-1]])+q*Matrix([[eigth,-eigth],[-eigth,-eigth]])
    newa = gauge(a,g,1)
    return newa

def cubic_surface(a):
    g0 = Matrix([[-6,-7,48],[-2,Rational(1,3),7],[1,0,1]])
    g1 = Matrix([[0,0,-Rational(11,9*27)],[Rational(1,9),0,-Rational(6,27)],[Rational(2,9*27),Rational(1,27*27),0]])
    g2 = Matrix([[0,0,Rational(166,27*2187)],[0,0,Rational(4,27*81)],\
                 [Rational(13,27*2187),Rational(2,27*729),0]])
    g = g0*(eye(3)+q*g1+q**2*g2)
    newa = gauge(a,g,3)
    return newa

def cubic_surface2(a):
    g = Matrix([[1,3*q**-1,0],[0,1,0],[0,0,1]])
    ginv = Matrix([[1,-3*q**-1,0],[0,1,0],[0,0,1]])
    newa = expand(ginv*a*g + ginv*diff(g,q))
    return newa

a = -q**(-2)*Matrix([[0,108,252],[1,9,36],[0,3,0]])+q**(-1)*Matrix([[0,0,0],[0,1,0],[0,0,2]])
print("Quantum connection for the cubic surface:")
print()
pprint(a)
print()
biga = cubic_surface(a)
print("After the first gauge transformation, up to O(q**2) error:")
print()
pprint(biga)
print()
print()
print("Further gauge transformation, up to O(1) error:")
print()
a = truncatea(cubic_surface2(biga),-1)
pprint(a)
print()
print(a)
