feat: 重构为 Go Web 脚手架——在线图片识别 Web 服务 + vision-cli 独立工具
- Gin + config.yaml + SQLite 后端骨架,JWT 登录 + 管理员初始化
- web/ 原生单页:拖拽多图上传 → Agnes 免费识别 → 结果复制
- HTTP API:/api/v1/vision/analyze 免鉴权 + 令牌桶限流,支持程序直调
- cmd/vision-cli 独立 CLI 工具(复用 service 层,不依赖 Web 服务)
- build.bat / build.sh 构建脚本,移除 mcp-server
- 统一响应格式 {code, message, data},CORS/JWT/限流中间件
Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
91
internal/config/config.go
Normal file
91
internal/config/config.go
Normal file
@@ -0,0 +1,91 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
// Config 是应用统一配置,对应 config.yaml。
|
||||
type Config struct {
|
||||
Server ServerConfig `yaml:"server"`
|
||||
Database DatabaseConfig `yaml:"database"`
|
||||
JWT JWTConfig `yaml:"jwt"`
|
||||
Upload UploadConfig `yaml:"upload"`
|
||||
AI AIConfig `yaml:"ai"`
|
||||
Env EnvConfig `yaml:"env"`
|
||||
}
|
||||
|
||||
type ServerConfig struct {
|
||||
Port int `yaml:"port"`
|
||||
}
|
||||
|
||||
type DatabaseConfig struct {
|
||||
Driver string `yaml:"driver"`
|
||||
Path string `yaml:"path"`
|
||||
AutoMigrate bool `yaml:"auto_migrate"`
|
||||
}
|
||||
|
||||
type JWTConfig struct {
|
||||
Secret string `yaml:"secret"`
|
||||
ExpireHours int `yaml:"expire_hours"`
|
||||
}
|
||||
|
||||
type UploadConfig struct {
|
||||
Path string `yaml:"path"`
|
||||
MaxSize int64 `yaml:"max_size"`
|
||||
}
|
||||
|
||||
type AIConfig struct {
|
||||
APIKey string `yaml:"api_key"`
|
||||
BaseURL string `yaml:"base_url"`
|
||||
Model string `yaml:"model"`
|
||||
}
|
||||
|
||||
type EnvConfig struct {
|
||||
Name string `yaml:"name"`
|
||||
}
|
||||
|
||||
// hardcodedAPIKey 内置默认免费 Key(与 config.yaml 一致),
|
||||
// 保证 CLI 工具在无配置文件时也能开箱即用。
|
||||
const hardcodedAPIKey = "sk-68FL9xXRrJhxcY5TCGxoaFlQ94oCBioIJhyGfeDMCkCvA0SV"
|
||||
|
||||
// DefaultConfig 返回开箱即用的默认配置。
|
||||
func DefaultConfig() *Config {
|
||||
return &Config{
|
||||
Server: ServerConfig{Port: 8080},
|
||||
Database: DatabaseConfig{Driver: "sqlite", Path: "./data/app.db", AutoMigrate: true},
|
||||
JWT: JWTConfig{Secret: "change-me-in-production", ExpireHours: 720},
|
||||
Upload: UploadConfig{Path: "uploads", MaxSize: 10 * 1024 * 1024},
|
||||
AI: AIConfig{
|
||||
APIKey: hardcodedAPIKey,
|
||||
Model: "agnes-2.0-flash",
|
||||
BaseURL: "https://apihub.agnes-ai.com/v1/chat/completions",
|
||||
},
|
||||
Env: EnvConfig{Name: "dev"},
|
||||
}
|
||||
}
|
||||
|
||||
// Load 从 yaml 文件加载配置;文件不存在时使用默认值(便于快速上手)。
|
||||
// 环境变量覆盖:AGNES_API_KEY → ai.api_key。
|
||||
func Load(path string) (*Config, error) {
|
||||
cfg := DefaultConfig()
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return cfg, nil
|
||||
}
|
||||
return nil, fmt.Errorf("读取 %s: %w", path, err)
|
||||
}
|
||||
if err := yaml.Unmarshal(data, cfg); err != nil {
|
||||
return nil, fmt.Errorf("解析 %s: %w", path, err)
|
||||
}
|
||||
|
||||
// 环境变量覆盖
|
||||
if key := os.Getenv("AGNES_API_KEY"); key != "" {
|
||||
cfg.AI.APIKey = key
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
48
internal/database/database.go
Normal file
48
internal/database/database.go
Normal file
@@ -0,0 +1,48 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
|
||||
"vision-tool/internal/config"
|
||||
"vision-tool/internal/model"
|
||||
)
|
||||
|
||||
// Init 初始化数据库连接并执行 AutoMigrate。
|
||||
// 使用 pure-go 的 glebarez/sqlite 驱动,CGO_ENABLED=0 下可编译运行。
|
||||
func Init(cfg *config.DatabaseConfig) (*gorm.DB, error) {
|
||||
var db *gorm.DB
|
||||
var err error
|
||||
|
||||
switch cfg.Driver {
|
||||
case "sqlite":
|
||||
fallthrough
|
||||
default:
|
||||
if dir := filepath.Dir(cfg.Path); dir != "" {
|
||||
if err := os.MkdirAll(dir, 0755); err != nil {
|
||||
return nil, fmt.Errorf("创建数据目录: %w", err)
|
||||
}
|
||||
}
|
||||
db, err = gorm.Open(sqlite.Open(cfg.Path), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Warn),
|
||||
})
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("连接数据库(%s): %w", cfg.Driver, err)
|
||||
}
|
||||
|
||||
// auto_migrate:自动建表/更新表结构,生产环境建议手动管理
|
||||
if cfg.AutoMigrate {
|
||||
if err := db.AutoMigrate(&model.User{}); err != nil {
|
||||
return nil, fmt.Errorf("自动建表: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return db, nil
|
||||
}
|
||||
88
internal/handler/auth_handler.go
Normal file
88
internal/handler/auth_handler.go
Normal file
@@ -0,0 +1,88 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"vision-tool/internal/model"
|
||||
"vision-tool/internal/service"
|
||||
)
|
||||
|
||||
// AuthHandler HTTP 处理器:管理员初始化 / 登录 / 用户信息。
|
||||
type AuthHandler struct {
|
||||
svc *service.AuthService
|
||||
}
|
||||
|
||||
func NewAuthHandler(svc *service.AuthService) *AuthHandler {
|
||||
return &AuthHandler{svc: svc}
|
||||
}
|
||||
|
||||
type initRequest struct {
|
||||
Username string `json:"username" binding:"required,min=2,max=32"`
|
||||
Password string `json:"password" binding:"required,min=6,max=64"`
|
||||
}
|
||||
|
||||
type loginRequest struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
|
||||
// Check GET /api/v1/admin/check — 管理员是否已初始化
|
||||
func (h *AuthHandler) Check(c *gin.Context) {
|
||||
initialized, err := h.svc.CheckAdmin()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, model.Err(model.CodeServerErr, err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, model.OK(gin.H{"initialized": initialized}))
|
||||
}
|
||||
|
||||
// Init POST /api/v1/admin/init — 初始化管理员(仅首次)
|
||||
func (h *AuthHandler) Init(c *gin.Context) {
|
||||
var req initRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, model.Err(model.CodeBadParam, "用户名 2-32 位,密码至少 6 位"))
|
||||
return
|
||||
}
|
||||
if err := h.svc.InitAdmin(req.Username, req.Password); err != nil {
|
||||
if errors.Is(err, service.ErrAdminExists) {
|
||||
c.JSON(http.StatusConflict, model.Err(model.CodeConflict, err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, model.Err(model.CodeServerErr, err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, model.OK(gin.H{"message": "管理员初始化成功"}))
|
||||
}
|
||||
|
||||
// Login POST /api/v1/auth/login — 登录获取 JWT token
|
||||
func (h *AuthHandler) Login(c *gin.Context) {
|
||||
var req loginRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, model.Err(model.CodeBadParam, "请输入用户名和密码"))
|
||||
return
|
||||
}
|
||||
token, err := h.svc.Login(req.Username, req.Password)
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrInvalidCredentials) {
|
||||
c.JSON(http.StatusUnauthorized, model.Err(model.CodeUnauthorized, err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, model.Err(model.CodeServerErr, err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, model.OK(gin.H{"token": token}))
|
||||
}
|
||||
|
||||
// Profile GET /api/v1/user/profile — 当前用户信息(JWT 保护)
|
||||
func (h *AuthHandler) Profile(c *gin.Context) {
|
||||
userID := c.GetUint("user_id")
|
||||
user, err := h.svc.Profile(userID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, model.Err(model.CodeServerErr, err.Error()))
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, model.OK(user))
|
||||
}
|
||||
124
internal/handler/vision_handler.go
Normal file
124
internal/handler/vision_handler.go
Normal file
@@ -0,0 +1,124 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"vision-tool/internal/config"
|
||||
"vision-tool/internal/model"
|
||||
"vision-tool/internal/service"
|
||||
)
|
||||
|
||||
// VisionHandler HTTP 处理器:在线图片识别。
|
||||
type VisionHandler struct {
|
||||
svc *service.VisionService
|
||||
upload config.UploadConfig
|
||||
}
|
||||
|
||||
func NewVisionHandler(svc *service.VisionService, upload config.UploadConfig) *VisionHandler {
|
||||
return &VisionHandler{svc: svc, upload: upload}
|
||||
}
|
||||
|
||||
// Analyze POST /api/v1/vision/analyze — 上传 1-N 张图片并调用免费视觉模型识别。
|
||||
//
|
||||
// multipart 表单字段:
|
||||
//
|
||||
// images 图片文件(可多个,支持 png/jpg/gif/webp/bmp/tiff)
|
||||
// prompt 可选,自定义分析提问(默认:结构化描述图片内容)
|
||||
//
|
||||
// 免鉴权 + 限流,方便 curl / Claude Code 直接调用。
|
||||
func (h *VisionHandler) Analyze(c *gin.Context) {
|
||||
form, err := c.MultipartForm()
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, model.Err(model.CodeBadParam, "请使用 multipart/form-data 上传图片"))
|
||||
return
|
||||
}
|
||||
|
||||
files := form.File["images"]
|
||||
if len(files) == 0 {
|
||||
c.JSON(http.StatusBadRequest, model.Err(model.CodeBadParam, "请至少上传一张图片 (images)"))
|
||||
return
|
||||
}
|
||||
// prompt 字段可选:未传时表单里没有该 key(切片为空),直接取 [0] 会越界 panic
|
||||
prompt := ""
|
||||
if len(form.Value["prompt"]) > 0 {
|
||||
prompt = strings.TrimSpace(form.Value["prompt"][0])
|
||||
}
|
||||
|
||||
// 保存上传图片到 uploads/,返回本地路径列表
|
||||
paths := make([]string, 0, len(files))
|
||||
urls := make([]string, 0, len(files))
|
||||
for _, fh := range files {
|
||||
path, url, err := h.saveImage(fh)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, model.Err(model.CodeBadParam, err.Error()))
|
||||
return
|
||||
}
|
||||
paths = append(paths, path)
|
||||
urls = append(urls, url)
|
||||
}
|
||||
|
||||
// 调用免费视觉模型(Web 端默认开启 thinking)
|
||||
result, err := h.svc.Analyze(paths, prompt, false)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, model.Err(model.CodeServerErr, "识别失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, model.OK(gin.H{
|
||||
"result": result.Content,
|
||||
"images": urls,
|
||||
"usage": result.Usage,
|
||||
}))
|
||||
}
|
||||
|
||||
// saveImage 校验并保存上传图片,返回磁盘路径与访问 URL。
|
||||
func (h *VisionHandler) saveImage(fh *multipart.FileHeader) (string, string, error) {
|
||||
ext := strings.ToLower(filepath.Ext(fh.Filename))
|
||||
allowed := map[string]bool{".png": true, ".jpg": true, ".jpeg": true, ".gif": true, ".webp": true, ".bmp": true, ".tiff": true, ".tif": true}
|
||||
if !allowed[ext] {
|
||||
return "", "", fmt.Errorf("不支持的图片格式: %s(支持 png/jpg/gif/webp/bmp/tiff)", fh.Filename)
|
||||
}
|
||||
if fh.Size > h.upload.MaxSize {
|
||||
return "", "", fmt.Errorf("图片超过大小限制: %s", fh.Filename)
|
||||
}
|
||||
|
||||
src, err := fh.Open()
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("读取上传文件: %w", err)
|
||||
}
|
||||
defer src.Close()
|
||||
|
||||
if err := os.MkdirAll(h.upload.Path, 0755); err != nil {
|
||||
return "", "", fmt.Errorf("创建上传目录: %w", err)
|
||||
}
|
||||
|
||||
name := randomName() + ext
|
||||
diskPath := filepath.Join(h.upload.Path, name)
|
||||
|
||||
dst, err := os.Create(diskPath)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("保存文件: %w", err)
|
||||
}
|
||||
defer dst.Close()
|
||||
if _, err := dst.ReadFrom(src); err != nil {
|
||||
return "", "", fmt.Errorf("写入文件: %w", err)
|
||||
}
|
||||
|
||||
return diskPath, "/uploads/" + name, nil
|
||||
}
|
||||
|
||||
// randomName 生成 16 字节 hex 随机文件名。
|
||||
func randomName() string {
|
||||
b := make([]byte, 16)
|
||||
rand.Read(b)
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
114
internal/middleware/middleware.go
Normal file
114
internal/middleware/middleware.go
Normal file
@@ -0,0 +1,114 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"math"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
|
||||
"vision-tool/internal/model"
|
||||
)
|
||||
|
||||
// CORS 跨域中间件(开发环境放开,生产环境应限制为具体域名)。
|
||||
func CORS() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
c.Header("Access-Control-Allow-Origin", "*")
|
||||
c.Header("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS")
|
||||
c.Header("Access-Control-Allow-Headers", "Content-Type, Authorization")
|
||||
if c.Request.Method == http.MethodOptions {
|
||||
c.AbortWithStatus(http.StatusNoContent)
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// JWT 鉴权
|
||||
// ============================================================================
|
||||
|
||||
// Claims JWT 载荷。
|
||||
type Claims struct {
|
||||
UserID uint `json:"uid"`
|
||||
Username string `json:"username"`
|
||||
Role string `json:"role"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
// JWTAuth 校验 Authorization: Bearer <token>,通过后将用户信息写入 context。
|
||||
func JWTAuth(secret string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
header := c.GetHeader("Authorization")
|
||||
tokenStr, ok := strings.CutPrefix(header, "Bearer ")
|
||||
if !ok || tokenStr == "" {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, model.Err(model.CodeUnauthorized, "未登录"))
|
||||
return
|
||||
}
|
||||
|
||||
claims := &Claims{}
|
||||
token, err := jwt.ParseWithClaims(tokenStr, claims, func(t *jwt.Token) (interface{}, error) {
|
||||
return []byte(secret), nil
|
||||
})
|
||||
if err != nil || !token.Valid {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, model.Err(model.CodeUnauthorized, "登录已过期,请重新登录"))
|
||||
return
|
||||
}
|
||||
|
||||
c.Set("user_id", claims.UserID)
|
||||
c.Set("username", claims.Username)
|
||||
c.Set("role", claims.Role)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 内存令牌桶限流(全局限流,防匿名刷接口)
|
||||
// ============================================================================
|
||||
|
||||
// RateLimiter 基于时间的令牌桶。
|
||||
type RateLimiter struct {
|
||||
mu sync.Mutex
|
||||
rate float64 // 每秒补充令牌数
|
||||
burst float64 // 桶容量
|
||||
tokens float64
|
||||
lastTime time.Time
|
||||
}
|
||||
|
||||
func NewRateLimiter(rps, burst int) *RateLimiter {
|
||||
return &RateLimiter{
|
||||
rate: float64(rps),
|
||||
burst: float64(burst),
|
||||
tokens: float64(burst),
|
||||
lastTime: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
func (rl *RateLimiter) Allow() bool {
|
||||
rl.mu.Lock()
|
||||
defer rl.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
rl.tokens = math.Min(rl.burst, rl.tokens+now.Sub(rl.lastTime).Seconds()*rl.rate)
|
||||
rl.lastTime = now
|
||||
|
||||
if rl.tokens >= 1 {
|
||||
rl.tokens--
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Middleware 返回限流 gin 中间件。
|
||||
func (rl *RateLimiter) Middleware() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if !rl.Allow() {
|
||||
c.AbortWithStatusJSON(http.StatusTooManyRequests, model.Err(429, "请求过于频繁,请稍后再试"))
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
38
internal/model/model.go
Normal file
38
internal/model/model.go
Normal file
@@ -0,0 +1,38 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
// User 管理员账号。
|
||||
type User struct {
|
||||
ID uint `gorm:"primaryKey" json:"id"`
|
||||
Username string `gorm:"uniqueIndex;size:64" json:"username"`
|
||||
PasswordHash string `gorm:"size:255" json:"-"`
|
||||
Role string `gorm:"size:32" json:"role"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// Response 统一响应格式 {code, message, data}。
|
||||
type Response struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
Data interface{} `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
// 业务错误码
|
||||
const (
|
||||
CodeOK = 0
|
||||
CodeBadParam = 400
|
||||
CodeUnauthorized = 401
|
||||
CodeForbidden = 403
|
||||
CodeNotFound = 404
|
||||
CodeConflict = 409
|
||||
CodeServerErr = 500
|
||||
)
|
||||
|
||||
func OK(data interface{}) *Response {
|
||||
return &Response{Code: CodeOK, Message: "ok", Data: data}
|
||||
}
|
||||
|
||||
func Err(code int, msg string) *Response {
|
||||
return &Response{Code: code, Message: msg}
|
||||
}
|
||||
45
internal/repository/user_repo.go
Normal file
45
internal/repository/user_repo.go
Normal file
@@ -0,0 +1,45 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"vision-tool/internal/model"
|
||||
)
|
||||
|
||||
// UserRepository 数据访问层:用户表 CRUD。
|
||||
type UserRepository struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func NewUserRepository(db *gorm.DB) *UserRepository {
|
||||
return &UserRepository{db: db}
|
||||
}
|
||||
|
||||
// Count 返回用户总数(用于判断管理员是否已初始化)。
|
||||
func (r *UserRepository) Count() (int64, error) {
|
||||
var count int64
|
||||
err := r.db.Model(&model.User{}).Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
|
||||
func (r *UserRepository) Create(user *model.User) error {
|
||||
return r.db.Create(user).Error
|
||||
}
|
||||
|
||||
func (r *UserRepository) FindByUsername(username string) (*model.User, error) {
|
||||
var user model.User
|
||||
err := r.db.Where("username = ?", username).First(&user).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
func (r *UserRepository) FindByID(id uint) (*model.User, error) {
|
||||
var user model.User
|
||||
err := r.db.First(&user, id).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
101
internal/router/router.go
Normal file
101
internal/router/router.go
Normal file
@@ -0,0 +1,101 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"io/fs"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"vision-tool/internal/config"
|
||||
"vision-tool/internal/handler"
|
||||
"vision-tool/internal/middleware"
|
||||
"vision-tool/internal/model"
|
||||
"vision-tool/internal/repository"
|
||||
"vision-tool/internal/service"
|
||||
)
|
||||
|
||||
// Setup 构建路由,并完成分层依赖注入:
|
||||
// db → repository → service → handler → router。
|
||||
// webFS 为 //go:embed web 的嵌入文件系统,用于内嵌前端。
|
||||
func Setup(cfg *config.Config, db *gorm.DB, webFS fs.FS) *gin.Engine {
|
||||
r := gin.New()
|
||||
r.Use(gin.Logger(), gin.Recovery(), middleware.CORS())
|
||||
|
||||
// ---------- 分层依赖注入 ----------
|
||||
userRepo := repository.NewUserRepository(db)
|
||||
authSvc := service.NewAuthService(userRepo, cfg.JWT.Secret, time.Duration(cfg.JWT.ExpireHours)*time.Hour)
|
||||
visionSvc := service.NewVisionService(&cfg.AI)
|
||||
authH := handler.NewAuthHandler(authSvc)
|
||||
visionH := handler.NewVisionHandler(visionSvc, cfg.Upload)
|
||||
|
||||
// 识别接口免鉴权,用令牌桶限流防刷
|
||||
analyzeLimiter := middleware.NewRateLimiter(10, 20)
|
||||
|
||||
// ---------- API 路由 ----------
|
||||
api := r.Group("/api")
|
||||
api.GET("/health", func(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, model.OK(gin.H{"status": "ok"}))
|
||||
})
|
||||
|
||||
v1 := api.Group("/v1")
|
||||
{
|
||||
// 管理员初始化(方式 A:check + init)
|
||||
v1.GET("/admin/check", authH.Check)
|
||||
v1.POST("/admin/init", authH.Init)
|
||||
|
||||
// 认证
|
||||
v1.POST("/auth/login", authH.Login)
|
||||
v1.GET("/user/profile", middleware.JWTAuth(cfg.JWT.Secret), authH.Profile)
|
||||
|
||||
// 在线图片识别(免鉴权 + 限流)
|
||||
v1.POST("/vision/analyze", analyzeLimiter.Middleware(), visionH.Analyze)
|
||||
}
|
||||
|
||||
// ---------- 静态资源 ----------
|
||||
// 上传图片
|
||||
r.StaticFS("/uploads", http.Dir(cfg.Upload.Path))
|
||||
|
||||
// 内嵌前端:未匹配路由统一走 NoRoute
|
||||
// /api/* /uploads/* → 404;embed 内文件 → 直接返回;其余 → index.html(SPA 回退)
|
||||
assets, err := fs.Sub(webFS, "web")
|
||||
if err != nil {
|
||||
panic("web 目录嵌入失败: " + err.Error())
|
||||
}
|
||||
r.NoRoute(func(c *gin.Context) {
|
||||
path := c.Request.URL.Path
|
||||
if strings.HasPrefix(path, "/api/") || strings.HasPrefix(path, "/uploads/") {
|
||||
c.JSON(http.StatusNotFound, model.Err(model.CodeNotFound, "接口不存在"))
|
||||
return
|
||||
}
|
||||
if data, err := fs.ReadFile(assets, strings.TrimPrefix(path, "/")); err == nil {
|
||||
c.Data(http.StatusOK, mimeByExt(path), data)
|
||||
return
|
||||
}
|
||||
index, _ := fs.ReadFile(assets, "index.html")
|
||||
c.Data(http.StatusOK, "text/html; charset=utf-8", index)
|
||||
})
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
// mimeByExt 按扩展名返回静态资源 Content-Type。
|
||||
func mimeByExt(path string) string {
|
||||
switch filepath.Ext(path) {
|
||||
case ".css":
|
||||
return "text/css; charset=utf-8"
|
||||
case ".js":
|
||||
return "application/javascript; charset=utf-8"
|
||||
case ".svg":
|
||||
return "image/svg+xml"
|
||||
case ".png":
|
||||
return "image/png"
|
||||
case ".jpg", ".jpeg":
|
||||
return "image/jpeg"
|
||||
default:
|
||||
return "text/html; charset=utf-8"
|
||||
}
|
||||
}
|
||||
89
internal/service/auth_service.go
Normal file
89
internal/service/auth_service.go
Normal file
@@ -0,0 +1,89 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"vision-tool/internal/middleware"
|
||||
"vision-tool/internal/model"
|
||||
"vision-tool/internal/repository"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrAdminExists = errors.New("管理员已初始化,禁止重复初始化")
|
||||
ErrInvalidCredentials = errors.New("用户名或密码错误")
|
||||
)
|
||||
|
||||
// AuthService 认证业务:管理员初始化 / 登录 / 用户信息。
|
||||
type AuthService struct {
|
||||
repo *repository.UserRepository
|
||||
secret string
|
||||
expire time.Duration
|
||||
}
|
||||
|
||||
func NewAuthService(repo *repository.UserRepository, secret string, expire time.Duration) *AuthService {
|
||||
return &AuthService{repo: repo, secret: secret, expire: expire}
|
||||
}
|
||||
|
||||
// CheckAdmin 判断是否已存在管理员(方式 A:check + init)。
|
||||
func (s *AuthService) CheckAdmin() (bool, error) {
|
||||
count, err := s.repo.Count()
|
||||
return count > 0, err
|
||||
}
|
||||
|
||||
// InitAdmin 初始化管理员,已存在时拒绝。
|
||||
func (s *AuthService) InitAdmin(username, password string) error {
|
||||
initialized, err := s.CheckAdmin()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if initialized {
|
||||
return ErrAdminExists
|
||||
}
|
||||
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
user := &model.User{
|
||||
Username: username,
|
||||
PasswordHash: string(hash),
|
||||
Role: "admin",
|
||||
}
|
||||
return s.repo.Create(user)
|
||||
}
|
||||
|
||||
// Login 校验账号密码,签发 JWT token。
|
||||
func (s *AuthService) Login(username, password string) (string, error) {
|
||||
user, err := s.repo.FindByUsername(username)
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return "", ErrInvalidCredentials
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
if bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)) != nil {
|
||||
return "", ErrInvalidCredentials
|
||||
}
|
||||
|
||||
claims := middleware.Claims{
|
||||
UserID: user.ID,
|
||||
Username: user.Username,
|
||||
Role: user.Role,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(s.expire)),
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
},
|
||||
}
|
||||
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(s.secret))
|
||||
}
|
||||
|
||||
// Profile 按 ID 返回用户信息。
|
||||
func (s *AuthService) Profile(userID uint) (*model.User, error) {
|
||||
return s.repo.FindByID(userID)
|
||||
}
|
||||
253
internal/service/vision_service.go
Normal file
253
internal/service/vision_service.go
Normal file
@@ -0,0 +1,253 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"vision-tool/internal/config"
|
||||
)
|
||||
|
||||
// ============================================================================
|
||||
// API Types (OpenAI-compatible format used by Agnes)
|
||||
// ============================================================================
|
||||
|
||||
type ChatMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content []MessagePart `json:"content"`
|
||||
}
|
||||
|
||||
type MessagePart struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text,omitempty"`
|
||||
ImageURL *ImageURL `json:"image_url,omitempty"`
|
||||
}
|
||||
|
||||
type ImageURL struct {
|
||||
URL string `json:"url"`
|
||||
Detail string `json:"detail,omitempty"` // "auto", "low", "high"
|
||||
}
|
||||
|
||||
type ChatRequest struct {
|
||||
Model string `json:"model"`
|
||||
Messages []ChatMessage `json:"messages"`
|
||||
ChatTemplateKwargs *ChatTemplateKwargs `json:"chat_template_kwargs,omitempty"`
|
||||
}
|
||||
|
||||
// ChatTemplateKwargs 开启 Agnes 的 thinking 模式(OpenAI 兼容格式)。
|
||||
type ChatTemplateKwargs struct {
|
||||
EnableThinking bool `json:"enable_thinking"`
|
||||
}
|
||||
|
||||
type ChatResponse struct {
|
||||
ID string `json:"id"`
|
||||
Choices []Choice `json:"choices"`
|
||||
Usage *Usage `json:"usage,omitempty"`
|
||||
Error *APIError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type APIError struct {
|
||||
Message string `json:"message"`
|
||||
Type string `json:"type"`
|
||||
Code string `json:"code"`
|
||||
}
|
||||
|
||||
type Choice struct {
|
||||
Index int `json:"index"`
|
||||
Message RespMessage `json:"message"`
|
||||
}
|
||||
|
||||
type RespMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
type Usage struct {
|
||||
PromptTokens int `json:"prompt_tokens"`
|
||||
CompletionTokens int `json:"completion_tokens"`
|
||||
TotalTokens int `json:"total_tokens"`
|
||||
}
|
||||
|
||||
// AnalyzeResult 识别结果。
|
||||
type AnalyzeResult struct {
|
||||
Content string
|
||||
Usage *Usage
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Vision Service
|
||||
// ============================================================================
|
||||
|
||||
// VisionService 调用 Agnes 免费视觉模型识别图片内容。
|
||||
type VisionService struct {
|
||||
apiKey string
|
||||
model string
|
||||
baseURL string
|
||||
httpClient *http.Client
|
||||
maxRetries int
|
||||
}
|
||||
|
||||
func NewVisionService(cfg *config.AIConfig) *VisionService {
|
||||
return &VisionService{
|
||||
apiKey: cfg.APIKey,
|
||||
model: cfg.Model,
|
||||
baseURL: cfg.BaseURL,
|
||||
maxRetries: 3,
|
||||
httpClient: &http.Client{Timeout: 120 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
// Analyze 分析一张或多张图片(本地路径或 http(s) URL 均可),返回模型文本与 token 用量。
|
||||
// noThinking 为 true 时关闭 thinking 模式(更快、更低负载)。
|
||||
func (s *VisionService) Analyze(images []string, prompt string, noThinking bool) (*AnalyzeResult, error) {
|
||||
if prompt == "" {
|
||||
if len(images) > 1 {
|
||||
prompt = "请详细描述这些图片的内容,比较它们之间的异同。请用中文回答。"
|
||||
} else {
|
||||
prompt = "请详细描述这张图片的内容。如果图片中有文字,请完整识别出来。请用中文回答。"
|
||||
}
|
||||
}
|
||||
|
||||
var parts []MessagePart
|
||||
for _, path := range images {
|
||||
content, err := s.buildImageContent(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取图片 %s: %w", path, err)
|
||||
}
|
||||
parts = append(parts, content)
|
||||
}
|
||||
parts = append(parts, MessagePart{Type: "text", Text: prompt})
|
||||
|
||||
req := ChatRequest{
|
||||
Model: s.model,
|
||||
Messages: []ChatMessage{
|
||||
{Role: "user", Content: parts},
|
||||
},
|
||||
ChatTemplateKwargs: &ChatTemplateKwargs{EnableThinking: !noThinking},
|
||||
}
|
||||
|
||||
resp, err := s.callWithRetry(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(resp.Choices) == 0 {
|
||||
return nil, fmt.Errorf("模型未返回内容")
|
||||
}
|
||||
return &AnalyzeResult{Content: resp.Choices[0].Message.Content, Usage: resp.Usage}, nil
|
||||
}
|
||||
|
||||
// buildImageContent 构建图片消息块:URL 直接使用,本地文件 base64 编码。
|
||||
func (s *VisionService) buildImageContent(path string) (MessagePart, error) {
|
||||
if strings.HasPrefix(path, "http://") || strings.HasPrefix(path, "https://") {
|
||||
return MessagePart{
|
||||
Type: "image_url",
|
||||
ImageURL: &ImageURL{URL: path},
|
||||
}, nil
|
||||
}
|
||||
|
||||
imgData, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return MessagePart{}, err
|
||||
}
|
||||
dataURL := fmt.Sprintf("data:%s;base64,%s", detectMimeType(path), base64.StdEncoding.EncodeToString(imgData))
|
||||
return MessagePart{
|
||||
Type: "image_url",
|
||||
ImageURL: &ImageURL{URL: dataURL},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// callWithRetry 调用 API,对 429/5xx/rate/busy 指数退避重试(2s → 5s → 8s → 11s)。
|
||||
func (s *VisionService) callWithRetry(req ChatRequest) (*ChatResponse, error) {
|
||||
var lastErr error
|
||||
for attempt := 0; attempt <= s.maxRetries; attempt++ {
|
||||
if attempt > 0 {
|
||||
delay := time.Duration(2+attempt*3) * time.Second
|
||||
fmt.Fprintf(os.Stderr, "Retrying in %v (attempt %d/%d)...\n", delay, attempt, s.maxRetries)
|
||||
time.Sleep(delay)
|
||||
}
|
||||
|
||||
resp, err := s.call(req)
|
||||
if err != nil {
|
||||
errStr := err.Error()
|
||||
if strings.Contains(errStr, "429") ||
|
||||
strings.Contains(errStr, "500") ||
|
||||
strings.Contains(errStr, "502") ||
|
||||
strings.Contains(errStr, "503") ||
|
||||
strings.Contains(errStr, "rate") ||
|
||||
strings.Contains(errStr, "busy") {
|
||||
lastErr = err
|
||||
continue
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
return nil, fmt.Errorf("重试 %d 次仍失败: %w", s.maxRetries, lastErr)
|
||||
}
|
||||
|
||||
func (s *VisionService) call(req ChatRequest) (*ChatResponse, error) {
|
||||
body, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("序列化请求: %w", err)
|
||||
}
|
||||
|
||||
httpReq, err := http.NewRequest(http.MethodPost, s.baseURL, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建请求: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Authorization", "Bearer "+s.apiKey)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := s.httpClient.Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("调用 API: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取响应: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
var chatResp ChatResponse
|
||||
if json.Unmarshal(respBody, &chatResp) == nil && chatResp.Error != nil {
|
||||
return nil, fmt.Errorf("API error (HTTP %d): [%s] %s",
|
||||
resp.StatusCode, chatResp.Error.Code, chatResp.Error.Message)
|
||||
}
|
||||
return nil, fmt.Errorf("API error (HTTP %d): %s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
var chatResp ChatResponse
|
||||
if err := json.Unmarshal(respBody, &chatResp); err != nil {
|
||||
return nil, fmt.Errorf("解析响应: %w\nRaw: %s", err, string(respBody))
|
||||
}
|
||||
return &chatResp, nil
|
||||
}
|
||||
|
||||
// detectMimeType 按扩展名判断图片 MIME 类型。
|
||||
func detectMimeType(path string) string {
|
||||
switch strings.ToLower(filepath.Ext(path)) {
|
||||
case ".jpg", ".jpeg":
|
||||
return "image/jpeg"
|
||||
case ".gif":
|
||||
return "image/gif"
|
||||
case ".webp":
|
||||
return "image/webp"
|
||||
case ".bmp":
|
||||
return "image/bmp"
|
||||
case ".tiff", ".tif":
|
||||
return "image/tiff"
|
||||
case ".png":
|
||||
fallthrough
|
||||
default:
|
||||
return "image/png"
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user