Tilemap/TilemapRenderObject.cs
using Sandbox.Rendering;
using System;
using System.Collections.Generic;

namespace Saandy.Tilemapper;

public sealed class TilemapRenderObject : SceneCustomObject
{
	public struct TileData
	{
		public Vector4 Position;
		public Vector4 UvRect; // xy = offset, zw = scale
	}

	public struct ChunkCoord : IEquatable<ChunkCoord>
	{
		public int X;
		public int Y;

		public ChunkCoord( int x, int y )
		{
			X = x;
			Y = y;
		}

		public bool Equals( ChunkCoord other ) => X == other.X && Y == other.Y;
		public override bool Equals( object obj ) => obj is ChunkCoord other && Equals( other );
		public override int GetHashCode() => HashCode.Combine( X, Y );
	}

	public sealed class TilemapChunk
	{
		public ChunkCoord Coord;
		public List<TilemapChunkBatch> Batches = new();
		public BBox Bounds;
		public int Revision;
		public bool HasContentHash;
		public int ContentHash;
	}

	public sealed class TilemapChunkBatch
	{
		public int LayerIndex;
		public ushort TilesetId;
		public GpuBuffer ArgsBuffer;
		public GpuBuffer<TileData> TileDataBuffer;
		public Material Material;

		public void Dispose()
		{
			TileDataBuffer?.Dispose();
			ArgsBuffer?.Dispose();

			TileDataBuffer = null;
			ArgsBuffer = null;
			Material = null;
		}
	}

	private readonly Material _tileMaterial;
	private readonly Model _tileModel;
	private readonly Vector3 _tileModelCenter;
	private readonly float _tileModelScale;

	private readonly Queue<TilemapChunk> _pendingDestroy = new();
	private readonly List<ChunkCoord> _removeChunks = new();
	private readonly Dictionary<(int LayerIndex, ushort TilesetId), List<TileData>> _buildGroups = new();
	private readonly uint[] _argsScratch = new uint[5];

	private int _tilemapRevision = -1;
	private TileMapAxis _tilemapAxis = (TileMapAxis)(-1);
	private int _tilemapLayerCount = -1;
	private float _tilemapLayerSpacing = -1.0f;
	private float _lastTileSize = -1.0f;
	private float _lastChunkSize = -1.0f;
	private int _lastChunkResolution = -1;
	private bool _disposed;

	public readonly Dictionary<ChunkCoord, TilemapChunk> ActiveChunks = new();

	public TileMap Tilemap { get; private set; }
	public CameraComponent CullingCamera { get; set; }
	public float TileSize { get; set; } = 1.0f;
	public float ChunkSize { get; set; } = 256.0f;
	public int ChunkResolution { get; set; } = 64;
	public int RenderRadius { get; set; } = 6;
	public TileMapAxis Axis { get; set; } = TileMapAxis.XZ;

	private float ChunkWorldSize => Math.Max( ChunkSize, Math.Max( TileSize, 0.0001f ) );

	public TilemapRenderObject( SceneWorld sceneWorld ) : base( sceneWorld )
	{
		_tileMaterial = Material.FromShader( "shaders/tilemap.shader" );
		_tileModel = LoadTileModel();

		BBox tileBounds = _tileModel.Bounds;
		_tileModelCenter = tileBounds.Center;

		float tileExtent = Math.Max( 0.0001f, Math.Max( tileBounds.Size.x, tileBounds.Size.z ) );
		_tileModelScale = 1.0f / tileExtent;

		// Keep the custom object itself from being culled before streamed chunks exist.
		// Individual chunks still use their own bounds when a culling camera is available.
		Bounds = BBox.FromPositionAndSize( Vector3.Zero, float.MaxValue );
	}

	private Model LoadTileModel()
	{
		return Model.Load( "models/tile/tile.vmdl" );
	}

	public void SetTilemap( TileMap tilemap )
	{
		if ( _disposed ) { return; }
		if ( Tilemap == tilemap ) { return; }

		Tilemap = tilemap;
		_tilemapRevision = -1;
		InvalidateAllChunks();
	}

	public void UpdateStreaming( Vector3 cameraPos )
	{
		if ( _disposed ) { return; }

		if ( Tilemap == null || !Tilemap.IsValid() )
		{
			DestroyMissingChunks();
			return;
		}

		TileSize = Math.Max( Tilemap.TileSize, 0.0001f );
		Axis = Tilemap.Axis;

		bool layoutChanged =
			_tilemapAxis != Tilemap.Axis ||
			_tilemapLayerCount != Tilemap.LayerCount ||
			Math.Abs( _tilemapLayerSpacing - Tilemap.LayerSpacing ) > 0.0001f ||
			Math.Abs( _lastTileSize - TileSize ) > 0.0001f ||
			Math.Abs( _lastChunkSize - ChunkSize ) > 0.0001f ||
			_lastChunkResolution != ChunkResolution;

		_tilemapRevision = Tilemap.Revision;
		_tilemapAxis = Tilemap.Axis;
		_tilemapLayerCount = Tilemap.LayerCount;
		_tilemapLayerSpacing = Tilemap.LayerSpacing;
		_lastTileSize = TileSize;
		_lastChunkSize = ChunkSize;
		_lastChunkResolution = ChunkResolution;

		if ( layoutChanged )
		{
			InvalidateAllChunks();
		}

		ChunkCoord cameraChunk = WorldToChunk( cameraPos );

		for ( int x = -RenderRadius; x <= RenderRadius; x++ )
		{
			for ( int y = -RenderRadius; y <= RenderRadius; y++ )
			{
				ChunkCoord coord = new( cameraChunk.X + x, cameraChunk.Y + y );

				if ( !ActiveChunks.TryGetValue( coord, out TilemapChunk chunk ) )
				{
					chunk = CreateChunk( coord );

					if ( chunk != null )
					{
						ActiveChunks.Add( coord, chunk );
					}
				}

				if ( chunk != null && chunk.Revision != Tilemap.Revision )
				{
					BuildChunk( chunk );
				}
			}
		}

		//
		// The visible set is always a square around cameraChunk, so there is no
		// reason to allocate a HashSet every frame just to discover removals.
		//
		_removeChunks.Clear();

		foreach ( var pair in ActiveChunks )
		{
			ChunkCoord coord = pair.Key;

			if ( Math.Abs( coord.X - cameraChunk.X ) > RenderRadius || Math.Abs( coord.Y - cameraChunk.Y ) > RenderRadius )
			{
				_removeChunks.Add( coord );
			}
		}

		foreach ( ChunkCoord coord in _removeChunks )
		{
			if ( !ActiveChunks.TryGetValue( coord, out TilemapChunk chunk ) ) { continue; }

			_pendingDestroy.Enqueue( chunk );
			ActiveChunks.Remove( coord );
		}

		_removeChunks.Clear();
		ProcessPendingDestroy();
	}

	private void InvalidateAllChunks()
	{
		foreach ( TilemapChunk chunk in ActiveChunks.Values )
		{
			chunk.Revision = -1;
		}
	}

	public void ProcessPendingDestroy()
	{
		while ( _pendingDestroy.Count > 0 )
		{
			TilemapChunk chunk = _pendingDestroy.Dequeue();
			DisposeChunk( chunk );
		}
	}

	private static void DisposeChunk( TilemapChunk chunk )
	{
		if ( chunk == null ) { return; }

		foreach ( TilemapChunkBatch batch in chunk.Batches )
		{
			batch.Dispose();
		}

		chunk.Batches.Clear();
	}

	public void Disable()
	{
		if ( ActiveChunks.Count > 0 )
		{
			foreach ( TilemapChunk chunk in ActiveChunks.Values )
			{
				_pendingDestroy.Enqueue( chunk );
			}

			ActiveChunks.Clear();
		}

		ProcessPendingDestroy();
	}

	/// <summary>
	/// Permanently releases this SceneCustomObject.
	///
	/// Disable() only releases streamed chunk resources. Shutdown() additionally
	/// removes the SceneCustomObject itself from the SceneWorld, which is required
	/// before TileMap drops its managed reference to this renderer.
	/// </summary>
	public void Shutdown()
	{
		if ( _disposed ) { return; }

		_disposed = true;

		Disable();

		Tilemap = null;
		CullingCamera = null;

		_removeChunks.Clear();

		foreach ( List<TileData> list in _buildGroups.Values )
		{
			list.Clear();
		}

		_buildGroups.Clear();

		Delete();
	}

	public override void RenderSceneObject()
	{
		if ( _disposed ) { return; }

		base.RenderSceneObject();

		if ( Tilemap == null || !Tilemap.IsValid() ) { return; }
		if ( CullingCamera == null ) { return; }

		foreach ( TilemapChunk chunk in ActiveChunks.Values )
		{
			RenderChunk( chunk );
		}
	}

	private TilemapChunk CreateChunk( ChunkCoord coord )
	{
		TilemapChunk chunk = new()
		{
			Coord = coord,
			Revision = -1,
			HasContentHash = false,
			ContentHash = 0
		};

		chunk.Bounds = BuildChunkBounds( coord );
		return chunk;
	}

	private void BuildChunk( TilemapChunk chunk )
	{
		if ( _disposed ) { return; }
		if ( Tilemap == null || !Tilemap.IsValid() ) { return; }
		if ( chunk == null ) { return; }

		//
		// Reuse all build lists. The previous implementation allocated a new
		// Dictionary plus one or more List<TileData>s for every chunk rebuild.
		//
		foreach ( List<TileData> list in _buildGroups.Values )
		{
			list.Clear();
		}

		int baseTileX = MathX.FloorToInt( (chunk.Coord.X * ChunkSize) / TileSize );
		int baseTileY = MathX.FloorToInt( (chunk.Coord.Y * ChunkSize) / TileSize );
		int maxTileX = baseTileX + ChunkResolution - 1;
		int maxTileY = baseTileY + ChunkResolution - 1;

		HashCode contentHash = new();

		// Bottom layers are farther from the camera. Topmost layer 0 is closest.
		// Each layer has its own tile data; layers never affect each other's masks.
		for ( int layerIndex = Tilemap.LayerCount - 1; layerIndex >= 0; layerIndex-- )
		{
			var layer = Tilemap.GetLayer( layerIndex );

			if ( layer == null || !layer.IsVisible ) { continue; }

			float layerOffset = Tilemap.GetLayerRenderNormalOffset( layerIndex );

			foreach ( Vector2Int cell in Tilemap.GetFilledCells( layerIndex ) )
			{
				if ( cell.x < baseTileX || cell.x > maxTileX || cell.y < baseTileY || cell.y > maxTileY ) { continue; }
				if ( !Tilemap.TryGetTile( layerIndex, cell.x, cell.y, out var tile ) ) { continue; }
				if ( !Tilemap.IsRenderableTile( tile ) ) { continue; }

				var tileset = Tilemap.GetTileset( tile.TilesetId );

				if ( tileset == null || !tileset.IsValid() ) { continue; }

				var key = (layerIndex, tile.TilesetId);

				if ( !_buildGroups.TryGetValue( key, out List<TileData> list ) )
				{
					list = new List<TileData>();
					_buildGroups[key] = list;
				}

				Vector3 worldPosition = Tilemap.CellToWorld( cell.x, cell.y, TileSize, layerOffset );
				Vector4 uvRect = tileset.GetUvRect( tile.TileId );

				TileData tileData = new()
				{
					Position = new Vector4( worldPosition.x, worldPosition.y, worldPosition.z, 1.0f ),
					UvRect = uvRect
				};

				list.Add( tileData );

				contentHash.Add( layerIndex );
				contentHash.Add( tile.TilesetId );
				contentHash.Add( tile.TileId );
				contentHash.Add( tileData.Position );
				contentHash.Add( tileData.UvRect );
				contentHash.Add( tileset.FilePath );
			}
		}

		int newContentHash = contentHash.ToHashCode();

		//
		// A TileMap revision is global, but most edits only touch one chunk.
		// Previously every active chunk destroyed/recreated its GPU buffers for
		// every paint operation. If this chunk's actual render data did not
		// change, just advance its revision and keep its existing GPU resources.
		//
		if ( chunk.HasContentHash && chunk.ContentHash == newContentHash )
		{
			chunk.Revision = Tilemap.Revision;
			chunk.Bounds = BuildChunkBounds( chunk.Coord );
			return;
		}

		DisposeChunk( chunk );

		foreach ( var pair in _buildGroups )
		{
			int layerIndex = pair.Key.LayerIndex;
			ushort tilesetId = pair.Key.TilesetId;
			List<TileData> data = pair.Value;

			if ( data.Count == 0 ) { continue; }

			var tileset = Tilemap.GetTileset( tilesetId );

			if ( tileset == null || !tileset.IsValid() ) { continue; }

			TilemapChunkBatch batch = new()
			{
				LayerIndex = layerIndex,
				TilesetId = tilesetId,
				Material = _tileMaterial.CreateCopy(),
				TileDataBuffer = new GpuBuffer<TileData>( data.Count ),
				ArgsBuffer = new GpuBuffer( 1, 5 * sizeof( uint ), GpuBuffer.UsageFlags.ByteAddress )
			};

			//
			// SetData needs a contiguous array here. This allocation now only
			// happens when this chunk's actual content changes, not whenever any
			// tile anywhere in the map changes.
			//
			batch.TileDataBuffer.SetData( data.ToArray() );

			FillArgs( _tileModel, (uint)data.Count );
			batch.ArgsBuffer.SetData( _argsScratch );

			if ( !string.IsNullOrWhiteSpace( tileset.FilePath ) )
			{
				Texture texture = Texture.Load( tileset.FilePath );
				batch.Material.Attributes.Set( "_TileTexture", texture );
			}

			chunk.Batches.Add( batch );
		}

		chunk.Batches.Sort( static ( a, b ) => b.LayerIndex.CompareTo( a.LayerIndex ) );

		chunk.HasContentHash = true;
		chunk.ContentHash = newContentHash;
		chunk.Revision = Tilemap.Revision;
		chunk.Bounds = BuildChunkBounds( chunk.Coord );
	}

	private void RenderChunk( TilemapChunk chunk )
	{
		if ( chunk == null ) { return; }

		if ( CullingCamera != null && !CullingCamera.GetFrustum().IsInside( chunk.Bounds, true ) )
		{
			return;
		}

		foreach ( TilemapChunkBatch batch in chunk.Batches )
		{
			if ( batch.ArgsBuffer == null || batch.TileDataBuffer == null || batch.Material == null ) { continue; }

			batch.Material.Attributes.Set( "_TileDataBuffer", batch.TileDataBuffer );
			batch.Material.Attributes.Set( "_TileSize", TileSize );
			batch.Material.Attributes.Set( "_TileModelCenter", _tileModelCenter );
			batch.Material.Attributes.Set( "_TileModelScale", _tileModelScale );
			batch.Material.Attributes.Set( "_TilePlaneAxisU", Tilemap.GetPlaneAxisU() );
			batch.Material.Attributes.Set( "_TilePlaneAxisV", Tilemap.GetPlaneAxisV() );
			batch.Material.Attributes.Set( "_TilePlaneNormal", Tilemap.GetPlaneNormal() );

			Graphics.DrawModelInstancedIndirect(
				_tileModel,
				batch.ArgsBuffer,
				0,
				batch.Material.Attributes
			);
		}
	}

	private BBox BuildChunkBounds( ChunkCoord coord )
	{
		float chunkWorldSize = ChunkWorldSize;
		float centerX = coord.X * chunkWorldSize + chunkWorldSize * 0.5f;
		float centerY = coord.Y * chunkWorldSize + chunkWorldSize * 0.5f;

		if ( Tilemap == null || !Tilemap.IsValid() )
		{
			return BBox.FromPositionAndSize( Vector3.Zero, chunkWorldSize );
		}

		return Tilemap.GetMapRectBounds( centerX, centerY, chunkWorldSize, chunkWorldSize, 10.0f + Tilemap.GetLayerRenderNormalThickness() );
	}

	private ChunkCoord WorldToChunk( Vector3 worldPos )
	{
		float chunkWorldSize = ChunkWorldSize;

		if ( Tilemap == null || !Tilemap.IsValid() )
		{
			return new ChunkCoord( 0, 0 );
		}

		Vector2 mapPosition = Tilemap.WorldToMap( worldPos );

		return new ChunkCoord(
			MathX.FloorToInt( mapPosition.x / chunkWorldSize ),
			MathX.FloorToInt( mapPosition.y / chunkWorldSize )
		);
	}

	private void FillArgs( Model model, uint instanceCount )
	{
		_argsScratch[0] = (uint)model.GetIndexCount( 0 );
		_argsScratch[1] = instanceCount;
		_argsScratch[2] = (uint)model.GetIndexStart( 0 );
		_argsScratch[3] = (uint)model.GetBaseVertex( 0 );
		_argsScratch[4] = 0;
	}

	private void DestroyMissingChunks()
	{
		if ( ActiveChunks.Count > 0 )
		{
			foreach ( TilemapChunk chunk in ActiveChunks.Values )
			{
				_pendingDestroy.Enqueue( chunk );
			}

			ActiveChunks.Clear();
		}

		ProcessPendingDestroy();
	}
}