import ee
import geemap
import cartopy.crs as ccrs
import matplotlib.pyplot as plt
from geemap import cartoee
import pyeepkg

# Initialize Earth Engine
ee.Initialize()

China = ee.FeatureCollection('projects/ee-pheobezhou/assets/China').geometry()
AI_index = ee.Image('projects/ee-pheobezhou/assets/ai_et0').clip(China).lte(6500).selfMask()
dem = ee.ImageCollection("JAXA/ALOS/AW3D30/V3_2").mosaic().select('DSM').mask(Vegcls)
Vegcls = ee.Image('projects/ee-pheobezhou/assets/vegcla2').clip(China).mask(AI_index).selfMask()


# Set plot style
plt.rcParams['font.family'] = 'Times New Roman'
plt.rcParams['font.size'] = 12

# Create figure and axis
fig = plt.figure(figsize=(5, 3))
ax = fig.add_subplot(1, 1, 1, projection=ccrs.PlateCarree())

# Define study region
bbox = [135, 24, 73, 56]  # [E, S, W, N]
plot_ext = ee.Geometry.Rectangle(bbox)

# Add DEM layer
palette = pyeepkg.get_palette_cptcity('DEM_print')
style_dem = {'palette': palette, 'min': 1000, 'max': 8000, 'opacity': 0.8}
cartoee.add_layer(ee_object=dem, ax=ax, vis_params=style_dem, region=bbox)

# Add grid lines
cartoee.add_gridlines(ax, interval=10, ytick_rotation=90, linestyle=":", draw_labels=False)

# Add colorbar
cartoee.add_colorbar(ax, vis_params=style_dem, loc='bottom',
                    orientation="horizontal", posOpts=[0.55, 0.25, 0.30, 0.03])

# Remove tick labels
ax.set_xticks([])
ax.set_yticks([])

# Save figure
fig.savefig(r"F:\re_results1\dem_mean.png", dpi=800, bbox_inches='tight')