99 lines
3.2 KiB
Go
99 lines
3.2 KiB
Go
package middleware
|
||
|
||
import (
|
||
"context"
|
||
"errors"
|
||
"net/http"
|
||
"strings"
|
||
|
||
"github.com/agentmail/gateway/internal/repo"
|
||
)
|
||
|
||
type contextKey string
|
||
|
||
const AgentNameKey contextKey = "agent_name"
|
||
|
||
// bearerToken 从 Authorization: Bearer <token> 取出令牌,缺失时返回空串。
|
||
func bearerToken(r *http.Request) string {
|
||
h := r.Header.Get("Authorization")
|
||
if h == "" {
|
||
return ""
|
||
}
|
||
const p = "Bearer "
|
||
if len(h) > len(p) && strings.EqualFold(h[:len(p)], p) {
|
||
return strings.TrimSpace(h[len(p):])
|
||
}
|
||
return ""
|
||
}
|
||
|
||
// BearerToken 导出给 handler 层用(注册接口不过中间件,需要自己取密钥)。
|
||
func BearerToken(r *http.Request) string { return bearerToken(r) }
|
||
|
||
// keyAuthError 把密钥校验错误翻译成对外文案。
|
||
// 「不存在」与「已使用/已过期」区分开:前两者是拿错了密钥,后者是密钥生命周期到了,
|
||
// 运维需要据此判断该重新签发还是该检查配置。
|
||
func keyAuthError(err error) string {
|
||
switch {
|
||
case errors.Is(err, repo.ErrKeyUsed):
|
||
return `{"error":"密钥已使用(一次性密钥只能用一次)"}`
|
||
case errors.Is(err, repo.ErrKeyExpired):
|
||
return `{"error":"密钥已过期"}`
|
||
default:
|
||
return `{"error":"密钥无效"}`
|
||
}
|
||
}
|
||
|
||
// AgentAuth 验证 Agent 身份,支持两种凭证:
|
||
//
|
||
// Authorization: Bearer <agent_key_token> —— 密钥认证(推荐)
|
||
// X-Agent-Name + X-Agent-Secret —— 旧的 name/secret 方式(兼容保留)
|
||
//
|
||
// 用户密钥(user_keys)不接受:两类密钥共享 token 命名空间但走各自的验证表,
|
||
// 因此用用户密钥调 Agent 接口只会得到「密钥无效」。
|
||
func AgentAuth(next http.Handler) http.Handler {
|
||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||
if token := bearerToken(r); token != "" {
|
||
agentName, err := repo.VerifyAgentKey(r.Context(), token)
|
||
if err != nil {
|
||
http.Error(w, keyAuthError(err), http.StatusUnauthorized)
|
||
return
|
||
}
|
||
if agentName == "" {
|
||
// 密钥有效但尚未绑定 Agent:注册接口会用请求里的 name 落定它,
|
||
// 其余接口无法确定调用者身份,只能拒。
|
||
http.Error(w, `{"error":"密钥尚未绑定 Agent,请先调用 /agent/register 完成注册"}`, http.StatusForbidden)
|
||
return
|
||
}
|
||
repo.HeartbeatAgent(r.Context(), agentName)
|
||
ctx := context.WithValue(r.Context(), AgentNameKey, agentName)
|
||
next.ServeHTTP(w, r.WithContext(ctx))
|
||
return
|
||
}
|
||
|
||
agentName := r.Header.Get("X-Agent-Name")
|
||
agentSecret := r.Header.Get("X-Agent-Secret")
|
||
if agentName == "" || agentSecret == "" {
|
||
http.Error(w, `{"error":"Missing Authorization: Bearer <key> or X-Agent-Name/X-Agent-Secret header"}`, http.StatusUnauthorized)
|
||
return
|
||
}
|
||
|
||
agent, err := repo.VerifyAgent(r.Context(), agentName, agentSecret)
|
||
if err != nil {
|
||
http.Error(w, `{"error":"Invalid credentials"}`, http.StatusUnauthorized)
|
||
return
|
||
}
|
||
|
||
repo.HeartbeatAgent(r.Context(), agent.Name)
|
||
ctx := context.WithValue(r.Context(), AgentNameKey, agent.Name)
|
||
next.ServeHTTP(w, r.WithContext(ctx))
|
||
})
|
||
}
|
||
|
||
// GetAgentName 从 context 中获取 agent_name
|
||
func GetAgentName(r *http.Request) string {
|
||
if v := r.Context().Value(AgentNameKey); v != nil {
|
||
return v.(string)
|
||
}
|
||
return ""
|
||
}
|