// Copyright (c) 2024-2026 Tencent Zhuque Lab. All rights reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
//     http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//
// Requirement: Any integration or derivative work must explicitly attribute
// Tencent Zhuque Lab (https://github.com/Tencent/AI-Infra-Guard) in its
// documentation or user interface, as detailed in the NOTICE file.

package middleware

import (
	"time"

	"github.com/gin-gonic/gin"
	"github.com/google/uuid"
	"trpc.group/trpc-go/trpc-go/log"
)

// TrpcMiddleware 创建trpc-go集成中间件
func TrpcMiddleware() gin.HandlerFunc {
	return func(c *gin.Context) {
		// 生成trace_id
		traceID := uuid.New().String()

		// 将trace_id放入gin context
		c.Set("trace_id", traceID)

		// 记录请求开始时间
		startTime := time.Now()

		// 记录请求开始日志
		log.Debugf("请求开始: trace_id=%s, method=%s, path=%s, client_ip=%s",
			traceID, c.Request.Method, c.FullPath(), getClientIP(c))

		// 继续处理请求
		c.Next()

		// 计算请求耗时
		duration := time.Since(startTime)

		// 记录请求结束日志
		log.Debugf("请求结束: trace_id=%s, method=%s, path=%s, status=%d, duration=%v",
			traceID, c.Request.Method, c.FullPath(), c.Writer.Status(), duration)

	}
}

// getClientIP 获取客户端真实IP
func getClientIP(c *gin.Context) string {
	// 优先从X-Real-IP获取
	if ip := c.GetHeader("X-Real-IP"); ip != "" {
		return ip
	}

	// 其次从X-Forwarded-For获取
	if ip := c.GetHeader("X-Forwarded-For"); ip != "" {
		return ip
	}

	// 最后从RemoteAddr获取
	return c.ClientIP()
}
