import numpy as np
from netCDF4 import Dataset as ncds
import matplotlib.pyplot as plt
import os
from copy import copy
import coarsen_input

def get_inputs(fct, topo_pct, wt_t, wt_u, wt_v,
               dp_tgt='',
               topo_sfx='',
               wavedrag_sfx='',
               sal_sfx='',
               porbar_sfx='',
               write=False):
    fn_topo = 'depth_p04arc'+str(fct)+topo_sfx+'.nc'
    fn_wavedrag = 'JSLH_p04arc'+str(fct)+wavedrag_sfx+'.nc'
    fn_sal = 'tpxo9a_sal_p04arc'+str(fct)+sal_sfx+'.nc'
    fn_porbar = 'depth_barrier_ave2_p04arc'+str(fct)+porbar_sfx+'.nc'

    coarsen_input.coarsen_hgrid(fct=fct,
                                write=write, tgt_dir=dp_tgt);
    coarsen_input.coarsen_topo(fct=fct, ocn_pct_crt=topo_pct, weight=wt_t,
                               write=write, tgt_dir=dp_tgt, tgt=fn_topo);
    coarsen_input.coarsen_wavedrag(fct=fct, weight=wt_t,
                                   write=write, tgt_dir=dp_tgt, tgt=fn_wavedrag);
    coarsen_input.coarsen_sal(fct=fct, weight=wt_t,
                              write=write, tgt_dir=dp_tgt, tgt=fn_sal);
    coarsen_input.coarsen_topo_barrier(fct=fct, weightu=wt_u, weightv=wt_v, ocn_pct_crt=0.0, oldmask=True,
                                       tgt_dir=dp_tgt, tgt=fn_porbar);

dp = '/Users/hewang/Research/subgrid_topo/input_p04c'
dp_src = os.path.join(dp, 'p04ar')

fn_grd_src = '/Users/hewang/Research/subgrid_topo/input_p04c/p04ar/hgrid_p04ar.nc'
area_h = ncds(fn_grd_src).variables['area'][:]

ny_h, nx_h = ncds(fn_grd_src).dimensions['ny'].size, ncds(fn_grd_src).dimensions['nx'].size

area_l = np.zeros((ny_h//2, nx_h//2))
for jj in range(2):
    for ii in range(2):
        area_l += area_h[jj::2, ii::2]

dx_h = ncds(fn_grd_src).variables['dx'][:]
dy_h = ncds(fn_grd_src).variables['dy'][:]
dx_l = dx_h[::2,::2] + dx_h[::2,1::2]
dy_l = dy_h[::2,::2] + dy_h[1::2,::2]

for fct in [2,3,6,9]:
    dp_tgt = os.path.join(dp, 'p04ar_c'+str(fct))
    get_inputs(fct, 0.5, wt_t=area_l, wt_u=dy_l, wt_v=dx_l,
               dp_tgt=dp_tgt,
               topo_sfx='_aw_moreocean_ocnpct0p5',
               wavedrag_sfx='_aw',
               sal_sfx='_aw',
               porbar_sfx='_aw_moreocean',
               write=True);