"""Bounded planar completion between real authored through-via terminals."""
import argparse,heapq,json,math,shutil,time
from pathlib import Path
import numpy as np
from shapely import contains_xy
from shapely.geometry import Point,LineString
from shapely.ops import unary_union
from kicad_tools.schema.pcb import PCB
from kicad_tools.validate.rules.clearance import _pad_polygon
from kicad_tools.geometry.copper import segment_copper_polygon
p=argparse.ArgumentParser();p.add_argument('source',type=Path);p.add_argument('ports',type=Path);p.add_argument('out',type=Path);p.add_argument('--nets',required=True);p.add_argument('--grid',type=float,default=.1);a=p.parse_args();a.out.mkdir(parents=True,exist_ok=True)
b=PCB.load(a.source);ports=json.loads(a.ports.read_text());results=[]
def obstacles(net,layer):
 shapes=[]
 for f in b.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 b.segments:
  if s.net_name!=net and s.layer==layer:
   shape=segment_copper_polygon(s.start,s.end,s.width)
   # Keep routed clock trunks three trace widths edge-to-edge from signals.
   # Authored short F package escapes are measured separately.
   if layer!='F.Cu' and (s.net_name=='SDCLK' or net=='SDCLK'):shape=shape.buffer(.39)
   shapes.append(shape)
 for v in b.vias:
  if v.net_name!=net:shapes.append(Point(v.position).buffer(v.size/2+(.39 if net=='SDCLK' or v.net_name=='SDCLK' else 0)))
 return unary_union(shapes).buffer(.15+.09+.005)
def findpath(start,end,obs):
 clear=lambda pts:not LineString(pts).intersects(obs)
 if clear([start,end]):return [start,end],0
 step=a.grid;xs=np.arange(.4,99.6,step);ys=np.arange(.4,79.6,step)
 # Half a grid diagonal guards the entire edge, not only its sampled endpoints.
 blocked=contains_xy(obs.buffer(step/math.sqrt(2)),xs[None,:],ys[:,None]);height,width=blocked.shape
 def xy(q):return (float(xs[q[0]]),float(ys[q[1]]))
 def seeds(pt):
  center=(round((pt[0]-.4)/step),round((pt[1]-.4)/step));answer=[]
  for dx in range(-4,5):
   for dy in range(-4,5):
    q=(center[0]+dx,center[1]+dy)
    if 0<=q[0]<width and 0<=q[1]<height and not blocked[q[1],q[0]] and clear([pt,xy(q)]):answer.append(q)
  return answer
 starts=seeds(start);goals=set(seeds(end));heap=[];dist={};previous={};visited=set();deadline=time.monotonic()+30
 for q in starts: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)<500000:
  _,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 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 not(0<=v[0]<width and 0<=v[1]<height) or blocked[v[1],v[0]] or v in visited:continue
   cost=dist[q]+step*math.hypot(dx,dy)
   if cost<dist.get(v,float('inf')):dist[v]=cost;previous[v]=q;heapq.heappush(heap,(cost+math.dist(xy(v),end),v))
 if found is None:return None,len(visited)
 points=[end,xy(found)];q=found
 while q in previous:q=previous[q];points.append(xy(q))
 points.append(start);points.reverse();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)
 return compact,len(visited)
for name in a.nets.split(','):
 endpoints=[tuple(v['port']) for v in ports if v['net']==name]
 if len(endpoints)!=2:print('SKIP endpoint count',name,len(endpoints),flush=True);continue
 layers=['In2.Cu','F.Cu'] if name.startswith('DQ') or name in ['LDQM','UDQM'] else ['In3.Cu','B.Cu']
 for layer in layers:
  start=time.monotonic();path,expanded=findpath(*endpoints,obstacles(name,layer));result={'net':name,'layer':layer,'routed':path is not None,'expanded':expanded,'seconds':round(time.monotonic()-start,3)};results.append(result);print(json.dumps(result),flush=True)
  if path:
   for u,v in zip(path,path[1:]):
    if math.dist(u,v)>1e-7:b.add_trace(u,v,width=.18,layer=layer,net=name)
   b.save(a.out/'physical.kicad_pcb');shutil.copy2(a.source.with_suffix('.kicad_pro'),a.out/'physical.kicad_pro');break
 (a.out/'results.json').write_text(json.dumps(results,indent=2))
