from pathlib import Path
import math,json,shutil,uuid,heapq
from shapely.geometry import Point,LineString
from shapely.ops import unary_union
from shapely.prepared import prep
from kicad_tools.schema.pcb import PCB
from kicad_tools.sexp import parse_string
from kicad_tools.validate.rules.clearance import _pad_polygon
from kicad_tools.geometry.copper import segment_copper_polygon
base=Path('/tmp/board07-sdram/aux-final');out=Path('/tmp/board07-sdram/dq14-detour-48');out.mkdir(exist_ok=True)
shutil.copy2(base/'sdram_demo.kicad_pro',out/'sdram_demo.kicad_pro');p=PCB.load(base/'sdram_demo.kicad_pcb')
def obstacles(net,layer):
 shapes=[]
 for f in p.footprints:
  for d in f.pads:
   if d.net_name!=net and (layer in d.layers or '*.Cu' in d.layers):
    q=_pad_polygon(d,f)
    if q is not None:shapes.append(q)
 for s in p.segments:
  if s.net_name!=net and s.layer==layer:shapes.append(segment_copper_polygon(s.start,s.end,s.width))
 for v in p.vias:
  if v.net_name!=net:shapes.append(Point(v.position).buffer(v.size/2))
 return unary_union(shapes).buffer(.153)
def clearpath(points,obs,width=.18):return not LineString(points).buffer(width/2).intersects(obs)
def addpath(points,net,layer):
 for a,b in zip(points,points[1:]):
  if math.dist(a,b)>.0001:p.add_trace(a,b,width=.18,layer=layer,net=net)
def padshapes():return unary_union([_pad_polygon(d,f) for f in p.footprints for d in f.pads if d.type=='smd'])
allpads=padshapes()
def viagood(q,net,obs):
 disc=Point(q).buffer(.25)
 if disc.distance(allpads)<.025 or any(disc.intersects(o) for o in obs.values()):return False
 return all(math.dist(q,v.position)>=.1+v.drill/2+.501 for v in p.vias)
def link(a,b,net,layer):
 obs=obstacles(net,layer);paths=[[a,b],[a,(a[0],b[1]),b],[a,(b[0],a[1]),b]]
 valid=[v for v in paths if clearpath(v,obs)]
 if valid:
  addpath(min(valid,key=lambda v:sum(math.dist(x,y) for x,y in zip(v,v[1:]))),net,layer);return
 # A bounded planar A* on existing obstacle polygons; all accepted segment
 # geometry is checked continuously, including the off-grid endpoint tails.
 step=.15;start=(round(a[0]/step),round(a[1]/step));end=(round(b[0]/step),round(b[1]/step));margin=10
 low=(int(min(a[0],b[0])/step)-int(margin/step),int(min(a[1],b[1])/step)-int(margin/step));high=(int(max(a[0],b[0])/step)+int(margin/step),int(max(a[1],b[1])/step)+int(margin/step))
 def xy(q):return(q[0]*step,q[1]*step)
 starts=[];ends=[]
 for dx in range(-2,3):
  for dy in range(-2,3):
   q=(start[0]+dx,start[1]+dy)
   if clearpath([a,xy(q)],obs):starts.append(q)
   q=(end[0]+dx,end[1]+dy)
   if clearpath([xy(q),b],obs):ends.append(q)
 goals=set(ends);heap=[];dist={};prev={};blocked={}
 for q in starts:dist[q]=math.dist(a,xy(q));heapq.heappush(heap,(dist[q]+math.dist(xy(q),b),q))
 visited=set();found=None
 while heap and len(visited)<150000:
  _,q=heapq.heappop(heap)
  if q in visited:continue
  visited.add(q)
  if q in goals:found=q;break
  for dx,dy in [(1,0),(-1,0),(0,1),(0,-1),(1,1),(-1,1),(1,-1),(-1,-1)]:
   v=(q[0]+dx,q[1]+dy)
   if v in visited or not(low[0]<=v[0]<=high[0] and low[1]<=v[1]<=high[1]):continue
   edge=(q,v)
   if edge not in blocked:blocked[edge]=not clearpath([xy(q),xy(v)],obs)
   if blocked[edge]:continue
   cost=dist[q]+step*math.hypot(dx,dy)
   if cost<dist.get(v,float('inf')):dist[v]=cost;prev[v]=q;heapq.heappush(heap,(cost+math.dist(xy(v),b),v))
 if found is None:raise RuntimeError(('No power route',net,a,b,layer,len(visited)))
 path=[b,xy(found)];q=found
 while q in prev:q=prev[q];path.append(xy(q))
 path.append(a);path=list(reversed(path));compact=[path[0]]
 for i in range(1,len(path)-1):
  u,v,w=compact[-1],path[i],path[i+1]
  if abs((v[0]-u[0])*(w[1]-v[1])-(v[1]-u[1])*(w[0]-v[0]))>1e-8:compact.append(v)
 compact.append(path[-1]);assert clearpath(compact,obs);addpath(compact,net,layer);print('A*',net,layer,len(visited),'segments',len(compact)-1,flush=True)


import numpy as np,time
from types import SimpleNamespace
from shapely import contains_xy
b=p
a=SimpleNamespace(grid=.1)
old_obstacles=obstacles
def obstacles(net,layer):
 obs=old_obstacles(net,layer).buffer(.092)
 clock=[segment_copper_polygon(x.start,x.end,x.width).buffer(.54+.09+.005) for x in b.segments if x.net_name=='SDCLK' and x.layer==layer]
 clockvias=[Point(v.position).buffer(v.size/2+.54+.09+.005) for v in b.vias if v.net_name=='SDCLK']
 return unary_union([obs,*clock,*clockvias])

def hybrid(net,start,end):
 layers=['In2.Cu','F.Cu'];obs=[obstacles(net,layer) for layer in layers];step=a.grid;xs=np.arange(.4,99.6,step);ys=np.arange(.4,79.6,step);width=len(xs);height=len(ys)
 blocked=np.stack([contains_xy(o.buffer(step/math.sqrt(2)),xs[None,:],ys[:,None]) for o in obs])
 # obstacles() already includes .245 mm clearance+trace radius. A .5 mm via
 # requires another .16 mm radius. Include foreign copper on EVERY layer.
 viaobs=unary_union([obstacles(net,layer).buffer(.16) for layer in ['F.Cu','In1.Cu','In2.Cu','In3.Cu','In4.Cu','B.Cu']])
 pads=unary_union([_pad_polygon(d,f) for f in b.footprints for d in f.pads if d.type=='smd']).buffer(.30)
 holes=unary_union([Point(v.position).buffer(.75) for v in b.vias])
 viablocked=contains_xy(unary_union([viaobs,pads,holes]),xs[None,:],ys[:,None])
 def xy(q):return(float(xs[q[0]]),float(ys[q[1]]))
 def clear(points,k):return not LineString(points).intersects(obs[k])
 def seeds(pt):
  cx,cy=round((pt[0]-.4)/step),round((pt[1]-.4)/step);answer=[]
  for k in range(2):
   for dx in range(-4,5):
    for dy in range(-4,5):
     q=(cx+dx,cy+dy,k)
     if 0<=q[0]<width and 0<=q[1]<height and not blocked[k,q[1],q[0]] and clear([pt,xy(q)],k):answer.append(q)
  return answer
 goals=set(seeds(end));heap=[];dist={};prev={};visited=set();deadline=time.monotonic()+90
 for q in seeds(start):dist[q]=math.dist(start,xy(q));heapq.heappush(heap,(dist[q]+math.dist(xy(q),end),q))
 found=None
 while heap and len(visited)<2000000:
  _,q=heapq.heappop(heap)
  if q in visited:continue
  visited.add(q)
  if len(visited)%10000==0 and time.monotonic()>deadline:break
  if q in goals:found=q;break
  for dx,dy,dk in [(1,0,0),(-1,0,0),(0,1,0),(0,-1,0),(1,1,0),(-1,1,0),(1,-1,0),(-1,-1,0),(0,0,1)]:
   v=(q[0]+dx,q[1]+dy,(q[2]+dk)%2 if dk else q[2])
   if dk and viablocked[q[1],q[0]]:continue
   if not(0<=v[0]<width and 0<=v[1]<height) or blocked[v[2],v[1],v[0]] or v in visited:continue
   cost=dist[q]+(3.0 if dk else step*math.hypot(dx,dy))
   if cost<dist.get(v,float('inf')):dist[v]=cost;prev[v]=q;heapq.heappush(heap,(cost+math.dist(xy(v),end),v))
 if found is None:return None,None,len(visited)
 states=[found];q=found
 while q in prev:q=prev[q];states.append(q)
 states.reverse();sections=[];vias=[];current=[start];k=states[0][2]
 for q in states:
  if q[2]!=k:
   sections.append((k,current));vias.append(xy(q));current=[xy(q)];k=q[2]
  else:current.append(xy(q))
 current.append(end);sections.append((k,current));traces=[]
 for k,points in sections:
  compact=[points[0]]
  for i in range(1,len(points)-1):
   u,v,w=compact[-1],points[i],points[i+1]
   if abs((v[0]-u[0])*(w[1]-v[1])-(v[1]-u[1])*(w[0]-v[0]))>1e-8:compact.append(v)
  compact.append(points[-1]);assert clear(compact,k)
  for u,v in zip(compact,compact[1:]):
   if math.dist(u,v)>1e-7:traces.append((u,v,layers[k]))
 assert all(not Point(q).intersects(viaobs) and not Point(q).intersects(pads) and not Point(q).intersects(holes) for q in vias)
 assert all(math.dist(u,v)>=.75 for i,u in enumerate(vias) for v in vias[i+1:])
 return traces,vias,len(visited)



source=PCB.load('boards/07-matchgroup-test/real_design/authored-source/sdram_demo.kicad_pcb')
protected={s.uuid for s in source.segments}|{v.uuid for v in source.vias}
number=b.get_net_by_name('DQ14').number
remove_ids={s.uuid for s in b.segments if s.net_name=='DQ14'}|{v.uuid for v in b.vias if v.net_name=='DQ14'}
remove_ids-=protected
for node in list(b._sexp.children):
 if node.name in ['segment','via'] and node.find_child('uuid').get_string(0) in remove_ids:b._sexp.remove(node)
b.save(out/'stripped.kicad_pcb');b=PCB.load(out/'stripped.kicad_pcb');p=b
endpoints=[v.position for v in b.vias if v.net_name=='DQ14'];print('ports',endpoints,flush=True)
orig_obs=obstacles
protected_leg=False
current_ports=[]
def obstacles(net,layer):
 obs=orig_obs(net,layer)
 if protected_leg:
  own=[LineString([x.start,x.end]).buffer(.72) for x in b.segments if x.net_name==net and x.layer==layer]
  own +=[Point(v.position).buffer(.88) for v in b.vias if v.net_name==net]
  region=unary_union(own).difference(unary_union([Point(q).buffer(1) for q in current_ports]))
  obs=unary_union([obs,region])
 return obs
waypoint=(85,48)
foreign={l.name:old_obstacles('DQ14',l.name) for l in b.copper_layers}
assert viagood(waypoint,'DQ14',foreign)
b.add_via(*waypoint,size=.5,drill=.2,net='DQ14')
for start,end in [(endpoints[0],waypoint),(waypoint,endpoints[1])]:
 current_ports=[start,end]
 traces,vias,count=hybrid('DQ14',start,end);print('leg',count,len(traces or []),len(vias or []),flush=True)
 assert traces
 for u,v,layer in traces:b.add_trace(u,v,width=.18,layer=layer,net='DQ14')
 for q in vias:b.add_via(*q,size=.5,drill=.2,net='DQ14')
 b.save(out/'sdram_demo.kicad_pcb');protected_leg=True
