From 01ecdc8c233c03783dcff41ad00e1220b2aa54d8 Mon Sep 17 00:00:00 2001 From: wonder Date: Fri, 26 Jun 2026 14:58:15 +0800 Subject: [PATCH] feat: initialize project skeleton with Go backend, database, and config - Go module with standard library + gorilla/sessions + mysql driver - Config: env-based configuration with .env file support - Database: MySQL connection with auto-migration for 6 tables - Seed data: 12 system tags with options, 8 system snippets - Auth: session cookie middleware - Handlers: auth, tags, snippets, builder, claude, dashboard, suggestions, settings - LLM: OpenAI-compatible client with 30s timeout - Docker: Dockerfile + docker-compose.yml - .env.example with all configuration options Co-Authored-By: Claude --- .env.example | 18 ++ Dockerfile | 18 ++ docker-compose.yml | 7 + go.mod | 10 + go.sum | 8 + internal/auth/middleware.go | 46 +++++ internal/config/config.go | 42 ++++ internal/db/db.go | 331 +++++++++++++++++++++++++++++++ internal/handlers/auth.go | 69 +++++++ internal/handlers/builder.go | 240 ++++++++++++++++++++++ internal/handlers/claude.go | 89 +++++++++ internal/handlers/dashboard.go | 161 +++++++++++++++ internal/handlers/settings.go | 54 +++++ internal/handlers/snippets.go | 115 +++++++++++ internal/handlers/suggestions.go | 59 ++++++ internal/handlers/tags.go | 182 +++++++++++++++++ internal/llm/client.go | 147 ++++++++++++++ internal/models/models.go | 91 +++++++++ main.go | 133 +++++++++++++ 19 files changed, 1820 insertions(+) create mode 100644 .env.example create mode 100644 Dockerfile create mode 100644 docker-compose.yml create mode 100644 go.mod create mode 100644 go.sum create mode 100644 internal/auth/middleware.go create mode 100644 internal/config/config.go create mode 100644 internal/db/db.go create mode 100644 internal/handlers/auth.go create mode 100644 internal/handlers/builder.go create mode 100644 internal/handlers/claude.go create mode 100644 internal/handlers/dashboard.go create mode 100644 internal/handlers/settings.go create mode 100644 internal/handlers/snippets.go create mode 100644 internal/handlers/suggestions.go create mode 100644 internal/handlers/tags.go create mode 100644 internal/llm/client.go create mode 100644 internal/models/models.go create mode 100644 main.go diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..faa4ce9 --- /dev/null +++ b/.env.example @@ -0,0 +1,18 @@ +# 服务配置 +SERVER_PORT=8080 +SESSION_SECRET=your-random-secret-key + +# 数据库(连接已有的远程 MySQL 实例) +DB_HOST=your-mysql-public-ip +DB_PORT=3306 +DB_USER=your-db-user +DB_PASSWORD=your-db-password +DB_NAME=prompt_generator + +# 认证 +AUTH_PASSWORD=your-access-password + +# LLM +LLM_API_BASE_URL=https://api.deepseek.com/v1 +LLM_API_KEY=sk-xxx +LLM_MODEL_NAME=deepseek-chat diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..73d0f18 --- /dev/null +++ b/Dockerfile @@ -0,0 +1,18 @@ +FROM golang:1.22-alpine AS builder + +WORKDIR /app +COPY go.mod go.sum ./ +RUN go mod download + +COPY . . +RUN CGO_ENABLED=0 GOOS=linux go build -o server . + +FROM alpine:latest +RUN apk --no-cache add ca-certificates + +WORKDIR /app +COPY --from=builder /app/server . +COPY --from=builder /app/frontend ./frontend + +EXPOSE 8080 +CMD ["./server"] diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..462f98c --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,7 @@ +services: + app: + build: . + ports: + - "8080:8080" + env_file: .env + restart: unless-stopped diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..9cb6fac --- /dev/null +++ b/go.mod @@ -0,0 +1,10 @@ +module prompt-generator + +go 1.26.3 + +require ( + filippo.io/edwards25519 v1.2.0 // indirect + github.com/go-sql-driver/mysql v1.10.0 // indirect + github.com/gorilla/securecookie v1.1.2 // indirect + github.com/gorilla/sessions v1.4.0 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..67b6d92 --- /dev/null +++ b/go.sum @@ -0,0 +1,8 @@ +filippo.io/edwards25519 v1.2.0 h1:crnVqOiS4jqYleHd9vaKZ+HKtHfllngJIiOpNpoJsjo= +filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q77mc= +github.com/go-sql-driver/mysql v1.10.0 h1:Q+1LV8DkHJvSYAdR83XzuhDaTykuDx0l6fkXxoWCWfw= +github.com/go-sql-driver/mysql v1.10.0/go.mod h1:M+cqaI7+xxXGG9swrdeUIoPG3Y3KCkF0pZej+SK+nWk= +github.com/gorilla/securecookie v1.1.2 h1:YCIWL56dvtr73r6715mJs5ZvhtnY73hBvEF8kXD8ePA= +github.com/gorilla/securecookie v1.1.2/go.mod h1:NfCASbcHqRSY+3a8tlWJwsQap2VX5pwzwo4h3eOamfo= +github.com/gorilla/sessions v1.4.0 h1:kpIYOp/oi6MG/p5PgxApU8srsSw9tuFbt46Lt7auzqQ= +github.com/gorilla/sessions v1.4.0/go.mod h1:FLWm50oby91+hl7p/wRxDth9bWSuk0qVL2emc7lT5ik= diff --git a/internal/auth/middleware.go b/internal/auth/middleware.go new file mode 100644 index 0000000..951f8f0 --- /dev/null +++ b/internal/auth/middleware.go @@ -0,0 +1,46 @@ +package auth + +import ( + "crypto/rand" + "encoding/hex" + "net/http" + "sync" + + "github.com/gorilla/sessions" +) + +var ( + Store *sessions.CookieStore + mu sync.Mutex +) + +func Init(secret string) { + Store = sessions.NewCookieStore([]byte(secret)) + Store.Options = &sessions.Options{ + Path: "/", + MaxAge: 86400 * 7, // 7 days + HttpOnly: true, + SameSite: http.SameSiteLaxMode, + } +} + +func GenerateToken() string { + b := make([]byte, 32) + rand.Read(b) + return hex.EncodeToString(b) +} + +// AuthMiddleware checks if the user is authenticated +func AuthMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + session, _ := Store.Get(r, "session") + auth, ok := session.Values["authenticated"].(bool) + if !ok || !auth { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusUnauthorized) + w.Write([]byte(`{"code":401,"message":"未登录"}`)) + return + } + next.ServeHTTP(w, r) + }) +} diff --git a/internal/config/config.go b/internal/config/config.go new file mode 100644 index 0000000..8bb7776 --- /dev/null +++ b/internal/config/config.go @@ -0,0 +1,42 @@ +package config + +import ( + "os" +) + +type Config struct { + ServerPort string + SessionSecret string + DBHost string + DBPort string + DBUser string + DBPassword string + DBName string + AuthPassword string + LLMAPIBaseURL string + LLMAPIKey string + LLMModelName string +} + +func Load() *Config { + return &Config{ + ServerPort: getEnv("SERVER_PORT", "8080"), + SessionSecret: getEnv("SESSION_SECRET", "default-secret-change-me"), + DBHost: getEnv("DB_HOST", "127.0.0.1"), + DBPort: getEnv("DB_PORT", "3306"), + DBUser: getEnv("DB_USER", "root"), + DBPassword: getEnv("DB_PASSWORD", ""), + DBName: getEnv("DB_NAME", "prompt_generator"), + AuthPassword: getEnv("AUTH_PASSWORD", "admin"), + LLMAPIBaseURL: getEnv("LLM_API_BASE_URL", "https://api.deepseek.com/v1"), + LLMAPIKey: getEnv("LLM_API_KEY", ""), + LLMModelName: getEnv("LLM_MODEL_NAME", "deepseek-chat"), + } +} + +func getEnv(key, fallback string) string { + if v := os.Getenv(key); v != "" { + return v + } + return fallback +} diff --git a/internal/db/db.go b/internal/db/db.go new file mode 100644 index 0000000..80708aa --- /dev/null +++ b/internal/db/db.go @@ -0,0 +1,331 @@ +package db + +import ( + "database/sql" + "fmt" + "log" + + _ "github.com/go-sql-driver/mysql" + + "prompt-generator/internal/config" +) + +var DB *sql.DB + +func Init(cfg *config.Config) error { + dsn := fmt.Sprintf("%s:%s@tcp(%s:%s)/%s?charset=utf8mb4&parseTime=true&loc=Local", + cfg.DBUser, cfg.DBPassword, cfg.DBHost, cfg.DBPort, cfg.DBName) + + var err error + DB, err = sql.Open("mysql", dsn) + if err != nil { + return fmt.Errorf("failed to open database: %w", err) + } + + if err = DB.Ping(); err != nil { + return fmt.Errorf("failed to ping database: %w", err) + } + + DB.SetMaxOpenConns(25) + DB.SetMaxIdleConns(5) + + log.Println("Database connected successfully") + return nil +} + +func AutoMigrate() error { + tables := []string{ + `CREATE TABLE IF NOT EXISTS tags ( + id bigint NOT NULL AUTO_INCREMENT, + name varchar(64) NOT NULL, + description varchar(255) NOT NULL DEFAULT '', + scope enum('system','personal') NOT NULL DEFAULT 'system', + created_at datetime NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (id), + KEY idx_scope (scope) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`, + + `CREATE TABLE IF NOT EXISTS tag_options ( + id bigint NOT NULL AUTO_INCREMENT, + tag_id bigint NOT NULL, + label varchar(128) NOT NULL, + constraint_text text NOT NULL, + sort_order int NOT NULL DEFAULT 0, + created_at datetime NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (id), + KEY idx_tag_id (tag_id) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`, + + `CREATE TABLE IF NOT EXISTS snippets ( + id bigint NOT NULL AUTO_INCREMENT, + name varchar(128) NOT NULL, + content text NOT NULL, + category varchar(64) NOT NULL DEFAULT '', + scope enum('system','personal') NOT NULL DEFAULT 'system', + created_at datetime NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (id), + KEY idx_scope (scope) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`, + + `CREATE TABLE IF NOT EXISTS builder_sessions ( + id bigint NOT NULL AUTO_INCREMENT, + title varchar(255) NOT NULL DEFAULT '', + project_name varchar(255) NOT NULL DEFAULT '', + final_prompt text, + claude_session_id varchar(128) DEFAULT NULL, + created_at datetime NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at datetime NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP, + PRIMARY KEY (id), + KEY idx_project_name (project_name), + KEY idx_claude_session_id (claude_session_id) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`, + + `CREATE TABLE IF NOT EXISTS builder_session_tags ( + id bigint NOT NULL AUTO_INCREMENT, + builder_session_id bigint NOT NULL, + tag_id bigint NOT NULL, + tag_option_id bigint NOT NULL, + PRIMARY KEY (id), + KEY idx_session_id (builder_session_id) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`, + + `CREATE TABLE IF NOT EXISTS builder_session_snippets ( + id bigint NOT NULL AUTO_INCREMENT, + builder_session_id bigint NOT NULL, + snippet_id bigint NOT NULL, + sort_order int NOT NULL DEFAULT 0, + PRIMARY KEY (id), + KEY idx_session_id (builder_session_id) + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`, + } + + for _, ddl := range tables { + if _, err := DB.Exec(ddl); err != nil { + return fmt.Errorf("failed to create table: %w", err) + } + } + + log.Println("Database migration completed") + return nil +} + +func SeedData() error { + // Check if system tags already exist + var count int + err := DB.QueryRow("SELECT COUNT(*) FROM tags WHERE scope='system'").Scan(&count) + if err != nil { + return err + } + if count > 0 { + log.Println("Seed data already exists, skipping") + return nil + } + + tx, err := DB.Begin() + if err != nil { + return err + } + defer tx.Rollback() + + // Insert system tags + type tagSeed struct { + name string + desc string + options []struct { + label string + constraint string + } + } + + seeds := []tagSeed{ + { + name: "Git 规范", + desc: "控制 Git 操作范围和提交规范", + options: []struct{ label, constraint string }{ + {"标准 commit,禁止 push", "所有变更必须先 commit,禁止直接 push 到远程仓库。每次 commit 需要清晰的提交信息。"}, + {"标准 commit,允许 push", "变更完成后可 commit 并 push 到远程仓库。"}, + {"仅允许 commit,禁止 merge/rebase", "只允许 commit 操作,禁止 merge、rebase、force push 等危险操作。"}, + }, + }, + { + name: "环境背景", + desc: "提供运行环境上下文信息", + options: []struct{ label, constraint string }{ + {"服务端环境(Linux 生产)", "当前工作环境为 Linux 生产服务器,注意路径分隔符和命令兼容性。"}, + {"本地开发环境", "当前为本地开发环境,可自由使用开发工具和调试手段。"}, + {"CI/CD 环境", "当前在 CI/CD 流水线中运行,注意环境变量和权限限制。"}, + {"测试环境", "当前为测试环境,可以进行破坏性测试。"}, + }, + }, + { + name: "工具约束", + desc: "限制可用的工具和操作方式", + options: []struct{ label, constraint string }{ + {"必须使用 AskUserQuestion 提问", "遇到需要用户确认的问题时,必须使用 AskUserQuestion 工具,不要自行判断。"}, + {"优先使用专用工具而非 shell", "优先使用 Read/Write/Edit 等专用工具,避免直接使用 shell 命令处理文件。"}, + {"禁止使用 Agent 工具", "不要使用 Agent 工具派发子任务,在当前会话中完成所有工作。"}, + }, + }, + { + name: "提交策略", + desc: "控制 Git 提交的粒度和时机", + options: []struct{ label, constraint string }{ + {"每完成一个小任务 commit 一次", "每完成一个独立的子任务就立即 commit,保持提交粒度细小。"}, + {"全部完成后一次性 commit", "所有任务完成后统一 commit 一个大的变更。"}, + {"按功能模块分别 commit", "按功能模块划分,每个模块完成后分别 commit。"}, + }, + }, + { + name: "阅读策略", + desc: "信息获取的顺序和方式", + options: []struct{ label, constraint string }{ + {"先读 docs/ 了解架构再读代码", "先阅读项目文档和架构说明,理解整体设计后再阅读具体代码实现。"}, + {"直接读代码", "直接阅读源代码,通过代码理解项目结构。"}, + {"先读 README 和 CLAUDE.md", "先阅读 README.md 和 CLAUDE.md 了解项目规范和约束。"}, + }, + }, + { + name: "输出格式", + desc: "约束输出内容的格式", + options: []struct{ label, constraint string }{ + {"Markdown", "输出内容使用 Markdown 格式,包含适当的标题、列表和代码块。"}, + {"JSON", "输出内容使用 JSON 格式,确保数据结构清晰。"}, + {"纯文本", "输出纯文本内容,不使用任何格式标记。"}, + {"代码块优先", "优先使用代码块展示内容,减少文字说明。"}, + }, + }, + { + name: "安全约束", + desc: "操作安全边界", + options: []struct{ label, constraint string }{ + {"禁止访问外部网络", "不要发起任何网络请求,不访问外部 API 或下载资源。"}, + {"禁止修改系统文件", "不要修改 /etc、/usr 等系统目录下的文件。"}, + {"禁止执行危险命令", "不要执行 rm -rf、chmod 777、dd 等危险命令。"}, + }, + }, + { + name: "代码风格", + desc: "代码编写规范", + options: []struct{ label, constraint string }{ + {"遵循项目现有风格", "严格遵循项目中已有的代码风格和命名约定,保持一致性。"}, + {"严格 ESLint/Prettier", "代码必须符合 ESLint 和 Prettier 规则,不得有 lint 错误。"}, + {"自由风格", "不限制代码风格,以可读性为优先。"}, + }, + }, + { + name: "测试要求", + desc: "测试策略和要求", + options: []struct{ label, constraint string }{ + {"必须编写单元测试", "所有新增功能必须配套单元测试,确保测试覆盖关键逻辑。"}, + {"仅手动验证", "通过手动测试验证功能,不强制要求自动化测试。"}, + {"TDD 方式", "采用测试驱动开发,先写测试再写实现。"}, + }, + }, + { + name: "错误处理", + desc: "遇到异常情况的行为", + options: []struct{ label, constraint string }{ + {"遇到不确定先问用户", "遇到不确定的情况时,使用 AskUserQuestion 向用户确认后再继续。"}, + {"尽可能自行判断", "尽量自行分析和判断,减少对用户的打扰。"}, + {"遇错停止等待指示", "遇到任何错误立即停止,等待用户进一步指示。"}, + }, + }, + { + name: "上下文策略", + desc: "Token 使用和输出策略", + options: []struct{ label, constraint string }{ + {"节约上下文,精简输出", "输出尽量精简,只包含必要信息,减少 token 消耗。"}, + {"详细输出,不省略", "输出完整详细的信息,不省略任何内容。"}, + {"平衡模式", "在信息完整性和 token 效率之间取得平衡。"}, + }, + }, + { + name: "工作模式", + desc: "工作执行方式", + options: []struct{ label, constraint string }{ + {"规划优先,先出方案再执行", "先制定详细方案,获得用户确认后再动手执行。"}, + {"直接执行,边做边调", "直接开始执行,根据反馈及时调整方向。"}, + {"探索模式,先研究再动手", "先深入研究和探索,充分理解后再开始执行。"}, + }, + }, + } + + for _, s := range seeds { + res, err := tx.Exec("INSERT INTO tags (name, description, scope) VALUES (?, ?, 'system')", s.name, s.desc) + if err != nil { + return err + } + tagID, _ := res.LastInsertId() + + for i, opt := range s.options { + _, err := tx.Exec("INSERT INTO tag_options (tag_id, label, constraint_text, sort_order) VALUES (?, ?, ?, ?)", + tagID, opt.label, opt.constraint, i) + if err != nil { + return err + } + } + } + + // Insert system snippets + snippets := []struct { + name string + content string + category string + }{ + { + name: "项目架构概述", + content: "## 项目架构\n\n请在开始工作前了解项目整体架构:\n- 后端框架和技术栈\n- 前端框架和技术栈\n- 数据库和缓存方案\n- 部署和运维方式\n\n确保你的修改符合项目整体架构设计。", + category: "架构", + }, + { + name: "技术栈说明", + content: "## 技术栈\n\n请在此处填写项目使用的主要技术栈:\n- 语言:\n- 框架:\n- 数据库:\n- 缓存:\n- 消息队列:\n- 部署:", + category: "架构", + }, + { + name: "目录结构描述", + content: "## 目录结构\n\n项目主要目录结构:\n```\n/\n├── src/ # 源代码\n├── tests/ # 测试文件\n├── docs/ # 文档\n├── config/ # 配置文件\n└── scripts/ # 脚本工具\n```\n\n请确保新增文件放在正确的目录下。", + category: "架构", + }, + { + name: "编码规范要求", + content: "## 编码规范\n\n1. 变量和函数命名使用 camelCase\n2. 常量使用 UPPER_SNAKE_CASE\n3. 类名使用 PascalCase\n4. 每个函数不超过 50 行\n5. 复杂逻辑必须添加注释\n6. 错误处理不能忽略", + category: "规范", + }, + { + name: "测试框架配置", + content: "## 测试要求\n\n- 使用项目现有的测试框架\n- 测试文件命名:xxx_test.go / xxx.test.ts\n- 测试覆盖率要求:核心逻辑 > 80%\n- 每个测试用例独立,不依赖执行顺序\n- Mock 外部依赖,不依赖真实服务", + category: "测试", + }, + { + name: "部署流程说明", + content: "## 部署流程\n\n1. 代码合并到 main 分支\n2. CI 自动运行测试\n3. 构建 Docker 镜像\n4. 部署到 staging 环境验证\n5. 手动确认后部署到生产环境\n\n注意:不要修改 CI/CD 配置文件,除非明确要求。", + category: "运维", + }, + { + name: "PR 规范", + content: "## PR 规范\n\n- PR 标题简洁明了,说明变更内容\n- 描述中包含:变更原因、主要改动、测试方式\n- 单个 PR 不超过 500 行变更\n- 必须通过 CI 检查\n- 至少一人 Code Review 后合并", + category: "规范", + }, + { + name: "Code Review 要点", + content: "## Code Review 要点\n\n- 代码逻辑是否正确\n- 是否有边界条件未处理\n- 是否有安全隐患\n- 命名是否清晰\n- 是否有冗余代码\n- 测试是否充分\n- 性能是否有问题", + category: "规范", + }, + } + + for _, s := range snippets { + _, err := tx.Exec("INSERT INTO snippets (name, content, category, scope) VALUES (?, ?, ?, 'system')", + s.name, s.content, s.category) + if err != nil { + return err + } + } + + if err := tx.Commit(); err != nil { + return err + } + + log.Println("Seed data inserted successfully") + return nil +} diff --git a/internal/handlers/auth.go b/internal/handlers/auth.go new file mode 100644 index 0000000..3cdab76 --- /dev/null +++ b/internal/handlers/auth.go @@ -0,0 +1,69 @@ +package handlers + +import ( + "encoding/json" + "net/http" + + "prompt-generator/internal/auth" + "prompt-generator/internal/config" + "prompt-generator/internal/models" +) + +var cfg *config.Config + +func Init(c *config.Config) { + cfg = c +} + +func writeJSON(w http.ResponseWriter, code int, resp models.APIResponse) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(code) + json.NewEncoder(w).Encode(resp) +} + +func success(w http.ResponseWriter, data interface{}) { + writeJSON(w, http.StatusOK, models.APIResponse{Code: 0, Message: "success", Data: data}) +} + +func fail(w http.ResponseWriter, httpCode int, msg string) { + writeJSON(w, httpCode, models.APIResponse{Code: httpCode, Message: msg}) +} + +func Login(w http.ResponseWriter, r *http.Request) { + var req struct { + Password string `json:"password"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + fail(w, 400, "请求格式错误") + return + } + + if req.Password != cfg.AuthPassword { + fail(w, 401, "密码错误") + return + } + + session, _ := auth.Store.Get(r, "session") + session.Values["authenticated"] = true + session.Save(r, w) + + success(w, nil) +} + +func Logout(w http.ResponseWriter, r *http.Request) { + session, _ := auth.Store.Get(r, "session") + session.Values["authenticated"] = false + session.Options.MaxAge = -1 + session.Save(r, w) + success(w, nil) +} + +func AuthCheck(w http.ResponseWriter, r *http.Request) { + session, _ := auth.Store.Get(r, "session") + authed, ok := session.Values["authenticated"].(bool) + if !ok || !authed { + fail(w, 401, "未登录") + return + } + success(w, map[string]bool{"authenticated": true}) +} diff --git a/internal/handlers/builder.go b/internal/handlers/builder.go new file mode 100644 index 0000000..bab6b88 --- /dev/null +++ b/internal/handlers/builder.go @@ -0,0 +1,240 @@ +package handlers + +import ( + "encoding/json" + "fmt" + "net/http" + "strconv" + "strings" + + "prompt-generator/internal/db" + "prompt-generator/internal/models" +) + +func GetBuilderSessions(w http.ResponseWriter, r *http.Request) { + rows, err := db.DB.Query("SELECT id, title, project_name, final_prompt, claude_session_id, created_at, updated_at FROM builder_sessions ORDER BY updated_at DESC") + if err != nil { + fail(w, 500, "查询会话列表失败") + return + } + defer rows.Close() + + var sessions []models.BuilderSession + for rows.Next() { + var s models.BuilderSession + if err := rows.Scan(&s.ID, &s.Title, &s.ProjectName, &s.FinalPrompt, &s.ClaudeSessionID, &s.CreatedAt, &s.UpdatedAt); err != nil { + continue + } + sessions = append(sessions, s) + } + if sessions == nil { + sessions = []models.BuilderSession{} + } + success(w, sessions) +} + +func CreateBuilderSession(w http.ResponseWriter, r *http.Request) { + var req struct { + Title string `json:"title"` + ProjectName string `json:"project_name"` + FinalPrompt string `json:"final_prompt"` + ClaudeSessionID string `json:"claude_session_id"` + TagOptionIDs []struct { + TagID int64 `json:"tag_id"` + TagOptionID int64 `json:"tag_option_id"` + } `json:"tag_option_ids"` + SnippetIDs []struct { + SnippetID int64 `json:"snippet_id"` + SortOrder int `json:"sort_order"` + } `json:"snippet_ids"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + fail(w, 400, "请求格式错误") + return + } + + var claudeID *string + if req.ClaudeSessionID != "" { + claudeID = &req.ClaudeSessionID + } + + res, err := db.DB.Exec("INSERT INTO builder_sessions (title, project_name, final_prompt, claude_session_id) VALUES (?, ?, ?, ?)", + req.Title, req.ProjectName, req.FinalPrompt, claudeID) + if err != nil { + fail(w, 500, "创建会话失败") + return + } + + sessionID, _ := res.LastInsertId() + + // Insert tag associations + for _, t := range req.TagOptionIDs { + db.DB.Exec("INSERT INTO builder_session_tags (builder_session_id, tag_id, tag_option_id) VALUES (?, ?, ?)", + sessionID, t.TagID, t.TagOptionID) + } + + // Insert snippet associations + for _, s := range req.SnippetIDs { + db.DB.Exec("INSERT INTO builder_session_snippets (builder_session_id, snippet_id, sort_order) VALUES (?, ?, ?)", + sessionID, s.SnippetID, s.SortOrder) + } + + success(w, map[string]int64{"id": sessionID}) +} + +func GetBuilderSession(w http.ResponseWriter, r *http.Request) { + id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) + if err != nil { + fail(w, 400, "无效的 ID") + return + } + + var s models.BuilderSession + err = db.DB.QueryRow("SELECT id, title, project_name, final_prompt, claude_session_id, created_at, updated_at FROM builder_sessions WHERE id=?", id). + Scan(&s.ID, &s.Title, &s.ProjectName, &s.FinalPrompt, &s.ClaudeSessionID, &s.CreatedAt, &s.UpdatedAt) + if err != nil { + fail(w, 404, "会话不存在") + return + } + + // Load tags + tagRows, err := db.DB.Query("SELECT id, builder_session_id, tag_id, tag_option_id FROM builder_session_tags WHERE builder_session_id=?", id) + if err == nil { + for tagRows.Next() { + var t models.SessionTag + if err := tagRows.Scan(&t.ID, &t.BuilderSessionID, &t.TagID, &t.TagOptionID); err == nil { + s.Tags = append(s.Tags, t) + } + } + tagRows.Close() + } + + // Load snippets + snippetRows, err := db.DB.Query("SELECT id, builder_session_id, snippet_id, sort_order FROM builder_session_snippets WHERE builder_session_id=? ORDER BY sort_order", id) + if err == nil { + for snippetRows.Next() { + var sn models.SessionSnippet + if err := snippetRows.Scan(&sn.ID, &sn.BuilderSessionID, &sn.SnippetID, &sn.SortOrder); err == nil { + s.Snippets = append(s.Snippets, sn) + } + } + snippetRows.Close() + } + + if s.Tags == nil { + s.Tags = []models.SessionTag{} + } + if s.Snippets == nil { + s.Snippets = []models.SessionSnippet{} + } + + success(w, s) +} + +func UpdateBuilderSession(w http.ResponseWriter, r *http.Request) { + id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) + if err != nil { + fail(w, 400, "无效的 ID") + return + } + + var req struct { + Title string `json:"title"` + ProjectName string `json:"project_name"` + FinalPrompt string `json:"final_prompt"` + ClaudeSessionID string `json:"claude_session_id"` + TagOptionIDs []struct { + TagID int64 `json:"tag_id"` + TagOptionID int64 `json:"tag_option_id"` + } `json:"tag_option_ids"` + SnippetIDs []struct { + SnippetID int64 `json:"snippet_id"` + SortOrder int `json:"sort_order"` + } `json:"snippet_ids"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + fail(w, 400, "请求格式错误") + return + } + + var claudeID *string + if req.ClaudeSessionID != "" { + claudeID = &req.ClaudeSessionID + } + + _, err = db.DB.Exec("UPDATE builder_sessions SET title=?, project_name=?, final_prompt=?, claude_session_id=? WHERE id=?", + req.Title, req.ProjectName, req.FinalPrompt, claudeID, id) + if err != nil { + fail(w, 500, "更新会话失败") + return + } + + // Replace tags + db.DB.Exec("DELETE FROM builder_session_tags WHERE builder_session_id=?", id) + for _, t := range req.TagOptionIDs { + db.DB.Exec("INSERT INTO builder_session_tags (builder_session_id, tag_id, tag_option_id) VALUES (?, ?, ?)", + id, t.TagID, t.TagOptionID) + } + + // Replace snippets + db.DB.Exec("DELETE FROM builder_session_snippets WHERE builder_session_id=?", id) + for _, s := range req.SnippetIDs { + db.DB.Exec("INSERT INTO builder_session_snippets (builder_session_id, snippet_id, sort_order) VALUES (?, ?, ?)", + id, s.SnippetID, s.SortOrder) + } + + success(w, nil) +} + +func DeleteBuilderSession(w http.ResponseWriter, r *http.Request) { + id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) + if err != nil { + fail(w, 400, "无效的 ID") + return + } + + db.DB.Exec("DELETE FROM builder_session_tags WHERE builder_session_id=?", id) + db.DB.Exec("DELETE FROM builder_session_snippets WHERE builder_session_id=?", id) + db.DB.Exec("DELETE FROM builder_sessions WHERE id=?", id) + success(w, nil) +} + +// AssemblePrompt 组装最终 Prompt +func AssemblePrompt(projectName string, customContent string, tagOptions []models.TagOption, snippets []models.Snippet, suggestions []string) string { + var parts []string + + // Project context + if projectName != "" { + parts = append(parts, fmt.Sprintf("## 项目上下文\n\n项目名称: %s", projectName)) + } + + // Custom content + if customContent != "" { + parts = append(parts, fmt.Sprintf("## 用户自定义内容\n\n%s", customContent)) + } + + // Tag constraints + if len(tagOptions) > 0 { + var constraints []string + for _, opt := range tagOptions { + constraints = append(constraints, fmt.Sprintf("- %s", opt.ConstraintText)) + } + parts = append(parts, fmt.Sprintf("## 标签约束\n\n%s", strings.Join(constraints, "\n"))) + } + + // Snippets + if len(snippets) > 0 { + var snippetParts []string + for _, s := range snippets { + snippetParts = append(snippetParts, fmt.Sprintf("--- 片段: %s ---\n%s", s.Name, s.Content)) + } + parts = append(parts, fmt.Sprintf("## 片段内容\n\n%s", strings.Join(snippetParts, "\n\n"))) + } + + // Suggestions + if len(suggestions) > 0 { + parts = append(parts, fmt.Sprintf("## 智能建议补充\n\n%s", strings.Join(suggestions, "\n"))) + } + + return strings.Join(parts, "\n\n") +} diff --git a/internal/handlers/claude.go b/internal/handlers/claude.go new file mode 100644 index 0000000..8e3ee0c --- /dev/null +++ b/internal/handlers/claude.go @@ -0,0 +1,89 @@ +package handlers + +import ( + "net/http" + + "prompt-generator/internal/db" +) + +func GetClaudeSessions(w http.ResponseWriter, r *http.Request) { + projectName := r.URL.Query().Get("project_name") + if projectName == "" { + fail(w, 400, "project_name 参数必填") + return + } + + rows, err := db.DB.Query(` + SELECT session_id, COUNT(*) as prompt_count, MAX(created_at) as last_time + FROM prompts + WHERE project_name=? + GROUP BY session_id + ORDER BY last_time DESC + LIMIT 50 + `, projectName) + if err != nil { + fail(w, 500, "查询失败") + return + } + defer rows.Close() + + var sessions []map[string]interface{} + for rows.Next() { + var sessionID string + var count int + var lastTime string + if err := rows.Scan(&sessionID, &count, &lastTime); err != nil { + continue + } + sessions = append(sessions, map[string]interface{}{ + "session_id": sessionID, + "prompt_count": count, + "last_time": lastTime, + }) + } + if sessions == nil { + sessions = []map[string]interface{}{} + } + success(w, sessions) +} + +func GetClaudePrompts(w http.ResponseWriter, r *http.Request) { + sessionID := r.PathValue("session_id") + if sessionID == "" { + fail(w, 400, "session_id 必填") + return + } + + rows, err := db.DB.Query(` + SELECT id, session_id, project_name, prompt, created_at + FROM prompts + WHERE session_id=? + ORDER BY created_at DESC + LIMIT 10 + `, sessionID) + if err != nil { + fail(w, 500, "查询失败") + return + } + defer rows.Close() + + var prompts []map[string]interface{} + for rows.Next() { + var id int64 + var sid, pname, prompt, createdAt string + if err := rows.Scan(&id, &sid, &pname, &prompt, &createdAt); err != nil { + continue + } + prompts = append(prompts, map[string]interface{}{ + "id": id, + "session_id": sid, + "project": pname, + "prompt": prompt, + "created_at": createdAt, + }) + } + if prompts == nil { + prompts = []map[string]interface{}{} + } + success(w, prompts) +} diff --git a/internal/handlers/dashboard.go b/internal/handlers/dashboard.go new file mode 100644 index 0000000..2c99443 --- /dev/null +++ b/internal/handlers/dashboard.go @@ -0,0 +1,161 @@ +package handlers + +import ( + "net/http" + "strconv" + + "prompt-generator/internal/db" +) + +func GetDashboardProjects(w http.ResponseWriter, r *http.Request) { + rows, err := db.DB.Query(` + SELECT project_name, COUNT(*) as session_count, MIN(created_at) as first_time, MAX(created_at) as last_time + FROM prompts + GROUP BY project_name + ORDER BY last_time DESC + `) + if err != nil { + fail(w, 500, "查询项目列表失败") + return + } + defer rows.Close() + + var projects []map[string]interface{} + for rows.Next() { + var name string + var count int + var firstTime, lastTime string + if err := rows.Scan(&name, &count, &firstTime, &lastTime); err != nil { + continue + } + if name == "" { + name = "(未命名项目)" + } + projects = append(projects, map[string]interface{}{ + "project_name": name, + "session_count": count, + "first_time": firstTime, + "last_time": lastTime, + }) + } + if projects == nil { + projects = []map[string]interface{}{} + } + success(w, projects) +} + +func GetDashboardSessions(w http.ResponseWriter, r *http.Request) { + projectName := r.PathValue("name") + if projectName == "" { + fail(w, 400, "项目名称必填") + return + } + + rows, err := db.DB.Query(` + SELECT session_id, COUNT(*) as prompt_count, MIN(created_at) as first_time, MAX(created_at) as last_time + FROM prompts + WHERE project_name=? + GROUP BY session_id + ORDER BY last_time DESC + LIMIT 100 + `, projectName) + if err != nil { + fail(w, 500, "查询会话列表失败") + return + } + defer rows.Close() + + var sessions []map[string]interface{} + for rows.Next() { + var sessionID string + var count int + var firstTime, lastTime string + if err := rows.Scan(&sessionID, &count, &firstTime, &lastTime); err != nil { + continue + } + sessions = append(sessions, map[string]interface{}{ + "session_id": sessionID, + "prompt_count": count, + "first_time": firstTime, + "last_time": lastTime, + }) + } + if sessions == nil { + sessions = []map[string]interface{}{} + } + success(w, sessions) +} + +func GetDashboardPrompts(w http.ResponseWriter, r *http.Request) { + sessionID := r.URL.Query().Get("session_id") + keyword := r.URL.Query().Get("keyword") + pageStr := r.URL.Query().Get("page") + pageSizeStr := r.URL.Query().Get("page_size") + + if sessionID == "" { + fail(w, 400, "session_id 参数必填") + return + } + + page := 1 + pageSize := 20 + if p, err := strconv.Atoi(pageStr); err == nil && p > 0 { + page = p + } + if ps, err := strconv.Atoi(pageSizeStr); err == nil && ps > 0 && ps <= 100 { + pageSize = ps + } + offset := (page - 1) * pageSize + + query := "SELECT id, session_id, project_name, prompt, created_at FROM prompts WHERE session_id=?" + countQuery := "SELECT COUNT(*) FROM prompts WHERE session_id=?" + args := []interface{}{sessionID} + countArgs := []interface{}{sessionID} + + if keyword != "" { + query += " AND prompt LIKE ?" + countQuery += " AND prompt LIKE ?" + kw := "%" + keyword + "%" + args = append(args, kw) + countArgs = append(countArgs, kw) + } + + var total int + db.DB.QueryRow(countQuery, countArgs...).Scan(&total) + + query += " ORDER BY created_at DESC LIMIT ? OFFSET ?" + args = append(args, pageSize, offset) + + rows, err := db.DB.Query(query, args...) + if err != nil { + fail(w, 500, "查询 Prompt 列表失败") + return + } + defer rows.Close() + + var prompts []map[string]interface{} + for rows.Next() { + var id int64 + var sid, pname, prompt, createdAt string + if err := rows.Scan(&id, &sid, &pname, &prompt, &createdAt); err != nil { + continue + } + prompts = append(prompts, map[string]interface{}{ + "id": id, + "session_id": sid, + "project": pname, + "prompt": prompt, + "created_at": createdAt, + }) + } + if prompts == nil { + prompts = []map[string]interface{}{} + } + + success(w, map[string]interface{}{ + "prompts": prompts, + "total": total, + "page": page, + "page_size": pageSize, + }) +} diff --git a/internal/handlers/settings.go b/internal/handlers/settings.go new file mode 100644 index 0000000..7a54138 --- /dev/null +++ b/internal/handlers/settings.go @@ -0,0 +1,54 @@ +package handlers + +import ( + "encoding/json" + "net/http" + + "prompt-generator/internal/db" + "prompt-generator/internal/models" +) + +func GetSettings(w http.ResponseWriter, r *http.Request) { + settings := models.Settings{ + LLMAPIBaseURL: cfg.LLMAPIBaseURL, + LLMModelName: cfg.LLMModelName, + } + success(w, settings) +} + +func UpdateSettings(w http.ResponseWriter, r *http.Request) { + var req struct { + LLMAPIBaseURL string `json:"llm_api_base_url"` + LLMAPIKey string `json:"llm_api_key"` + LLMModelName string `json:"llm_model_name"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + fail(w, 400, "请求格式错误") + return + } + + // Update in-memory config + if req.LLMAPIBaseURL != "" { + cfg.LLMAPIBaseURL = req.LLMAPIBaseURL + } + if req.LLMModelName != "" { + cfg.LLMModelName = req.LLMModelName + } + if req.LLMAPIKey != "" { + cfg.LLMAPIKey = req.LLMAPIKey + } + + // Persist to database (use a simple key-value table) + db.DB.Exec(`CREATE TABLE IF NOT EXISTS settings ( + key_name varchar(64) NOT NULL PRIMARY KEY, + value text NOT NULL + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4`) + + db.DB.Exec("REPLACE INTO settings (key_name, value) VALUES (?, ?)", "llm_api_base_url", cfg.LLMAPIBaseURL) + db.DB.Exec("REPLACE INTO settings (key_name, value) VALUES (?, ?)", "llm_model_name", cfg.LLMModelName) + if req.LLMAPIKey != "" { + db.DB.Exec("REPLACE INTO settings (key_name, value) VALUES (?, ?)", "llm_api_key", cfg.LLMAPIKey) + } + + success(w, nil) +} diff --git a/internal/handlers/snippets.go b/internal/handlers/snippets.go new file mode 100644 index 0000000..91df9bc --- /dev/null +++ b/internal/handlers/snippets.go @@ -0,0 +1,115 @@ +package handlers + +import ( + "encoding/json" + "net/http" + "strconv" + + "prompt-generator/internal/db" + "prompt-generator/internal/models" +) + +func GetSnippets(w http.ResponseWriter, r *http.Request) { + category := r.URL.Query().Get("category") + + query := "SELECT id, name, content, category, scope, created_at FROM snippets" + var args []interface{} + if category != "" { + query += " WHERE category=?" + args = append(args, category) + } + query += " ORDER BY scope, category, id" + + rows, err := db.DB.Query(query, args...) + if err != nil { + fail(w, 500, "查询片段失败") + return + } + defer rows.Close() + + var snippets []models.Snippet + for rows.Next() { + var s models.Snippet + if err := rows.Scan(&s.ID, &s.Name, &s.Content, &s.Category, &s.Scope, &s.CreatedAt); err != nil { + continue + } + snippets = append(snippets, s) + } + if snippets == nil { + snippets = []models.Snippet{} + } + success(w, snippets) +} + +func CreateSnippet(w http.ResponseWriter, r *http.Request) { + var req struct { + Name string `json:"name"` + Content string `json:"content"` + Category string `json:"category"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil || req.Name == "" { + fail(w, 400, "片段名称不能为空") + return + } + + res, err := db.DB.Exec("INSERT INTO snippets (name, content, category, scope) VALUES (?, ?, ?, 'personal')", + req.Name, req.Content, req.Category) + if err != nil { + fail(w, 500, "创建片段失败") + return + } + + id, _ := res.LastInsertId() + success(w, map[string]int64{"id": id}) +} + +func UpdateSnippet(w http.ResponseWriter, r *http.Request) { + id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) + if err != nil { + fail(w, 400, "无效的 ID") + return + } + + var scope string + db.DB.QueryRow("SELECT scope FROM snippets WHERE id=?", id).Scan(&scope) + if scope != "personal" { + fail(w, 403, "不能修改系统片段") + return + } + + var req struct { + Name string `json:"name"` + Content string `json:"content"` + Category string `json:"category"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + fail(w, 400, "请求格式错误") + return + } + + _, err = db.DB.Exec("UPDATE snippets SET name=?, content=?, category=? WHERE id=?", + req.Name, req.Content, req.Category, id) + if err != nil { + fail(w, 500, "更新片段失败") + return + } + success(w, nil) +} + +func DeleteSnippet(w http.ResponseWriter, r *http.Request) { + id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) + if err != nil { + fail(w, 400, "无效的 ID") + return + } + + var scope string + db.DB.QueryRow("SELECT scope FROM snippets WHERE id=?", id).Scan(&scope) + if scope != "personal" { + fail(w, 403, "不能删除系统片段") + return + } + + db.DB.Exec("DELETE FROM snippets WHERE id=?", id) + success(w, nil) +} diff --git a/internal/handlers/suggestions.go b/internal/handlers/suggestions.go new file mode 100644 index 0000000..004cf6f --- /dev/null +++ b/internal/handlers/suggestions.go @@ -0,0 +1,59 @@ +package handlers + +import ( + "encoding/json" + "net/http" + + "prompt-generator/internal/llm" + "prompt-generator/internal/models" +) + +func GetSuggestions(w http.ResponseWriter, r *http.Request) { + var req struct { + CurrentPrompt string `json:"current_prompt"` + ClaudeMD string `json:"claude_md"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + fail(w, 400, "请求格式错误") + return + } + + if req.CurrentPrompt == "" { + fail(w, 400, "当前 Prompt 不能为空") + return + } + + client := llm.NewClient(cfg.LLMAPIBaseURL, cfg.LLMAPIKey, cfg.LLMModelName) + suggestions, err := client.GetSuggestions(req.CurrentPrompt, req.ClaudeMD) + if err != nil { + fail(w, 500, "获取建议失败: "+err.Error()) + return + } + + success(w, suggestions) +} + +// parseSuggestions 尝试从 LLM 响应中解析建议列表 +func parseSuggestions(data string) []models.Suggestion { + // Try to find JSON array in the response + start := -1 + end := -1 + for i := 0; i < len(data); i++ { + if data[i] == '[' && start == -1 { + start = i + } + if data[i] == ']' { + end = i + 1 + } + } + + if start == -1 || end == -1 { + return nil + } + + var suggestions []models.Suggestion + if err := json.Unmarshal([]byte(data[start:end]), &suggestions); err != nil { + return nil + } + return suggestions +} diff --git a/internal/handlers/tags.go b/internal/handlers/tags.go new file mode 100644 index 0000000..b477026 --- /dev/null +++ b/internal/handlers/tags.go @@ -0,0 +1,182 @@ +package handlers + +import ( + "encoding/json" + "net/http" + "strconv" + + "prompt-generator/internal/db" + "prompt-generator/internal/models" +) + +func GetTags(w http.ResponseWriter, r *http.Request) { + rows, err := db.DB.Query("SELECT id, name, description, scope, created_at FROM tags ORDER BY scope, id") + if err != nil { + fail(w, 500, "查询标签失败") + return + } + defer rows.Close() + + var tags []models.Tag + for rows.Next() { + var t models.Tag + if err := rows.Scan(&t.ID, &t.Name, &t.Description, &t.Scope, &t.CreatedAt); err != nil { + continue + } + // Load options + optRows, err := db.DB.Query("SELECT id, tag_id, label, constraint_text, sort_order, created_at FROM tag_options WHERE tag_id=? ORDER BY sort_order", t.ID) + if err == nil { + for optRows.Next() { + var o models.TagOption + if err := optRows.Scan(&o.ID, &o.TagID, &o.Label, &o.ConstraintText, &o.SortOrder, &o.CreatedAt); err == nil { + t.Options = append(t.Options, o) + } + } + optRows.Close() + } + tags = append(tags, t) + } + + if tags == nil { + tags = []models.Tag{} + } + success(w, tags) +} + +func CreateTag(w http.ResponseWriter, r *http.Request) { + var req struct { + Name string `json:"name"` + Description string `json:"description"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil || req.Name == "" { + fail(w, 400, "标签名称不能为空") + return + } + + res, err := db.DB.Exec("INSERT INTO tags (name, description, scope) VALUES (?, ?, 'personal')", req.Name, req.Description) + if err != nil { + fail(w, 500, "创建标签失败") + return + } + + id, _ := res.LastInsertId() + success(w, map[string]int64{"id": id}) +} + +func UpdateTag(w http.ResponseWriter, r *http.Request) { + id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) + if err != nil { + fail(w, 400, "无效的 ID") + return + } + + // Check ownership + var scope string + db.DB.QueryRow("SELECT scope FROM tags WHERE id=?", id).Scan(&scope) + if scope != "personal" { + fail(w, 403, "不能修改系统标签") + return + } + + var req struct { + Name string `json:"name"` + Description string `json:"description"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + fail(w, 400, "请求格式错误") + return + } + + _, err = db.DB.Exec("UPDATE tags SET name=?, description=? WHERE id=?", req.Name, req.Description, id) + if err != nil { + fail(w, 500, "更新标签失败") + return + } + success(w, nil) +} + +func DeleteTag(w http.ResponseWriter, r *http.Request) { + id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) + if err != nil { + fail(w, 400, "无效的 ID") + return + } + + var scope string + db.DB.QueryRow("SELECT scope FROM tags WHERE id=?", id).Scan(&scope) + if scope != "personal" { + fail(w, 403, "不能删除系统标签") + return + } + + // Delete options first + db.DB.Exec("DELETE FROM tag_options WHERE tag_id=?", id) + db.DB.Exec("DELETE FROM tags WHERE id=?", id) + success(w, nil) +} + +func CreateTagOption(w http.ResponseWriter, r *http.Request) { + tagID, err := strconv.ParseInt(r.PathValue("id"), 10, 64) + if err != nil { + fail(w, 400, "无效的标签 ID") + return + } + + var req struct { + Label string `json:"label"` + ConstraintText string `json:"constraint_text"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil || req.Label == "" { + fail(w, 400, "选项名称不能为空") + return + } + + // Get max sort order + var maxOrder int + db.DB.QueryRow("SELECT COALESCE(MAX(sort_order),0) FROM tag_options WHERE tag_id=?", tagID).Scan(&maxOrder) + + res, err := db.DB.Exec("INSERT INTO tag_options (tag_id, label, constraint_text, sort_order) VALUES (?, ?, ?, ?)", + tagID, req.Label, req.ConstraintText, maxOrder+1) + if err != nil { + fail(w, 500, "创建选项失败") + return + } + + id, _ := res.LastInsertId() + success(w, map[string]int64{"id": id}) +} + +func UpdateTagOption(w http.ResponseWriter, r *http.Request) { + id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) + if err != nil { + fail(w, 400, "无效的 ID") + return + } + + var req struct { + Label string `json:"label"` + ConstraintText string `json:"constraint_text"` + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + fail(w, 400, "请求格式错误") + return + } + + _, err = db.DB.Exec("UPDATE tag_options SET label=?, constraint_text=? WHERE id=?", req.Label, req.ConstraintText, id) + if err != nil { + fail(w, 500, "更新选项失败") + return + } + success(w, nil) +} + +func DeleteTagOption(w http.ResponseWriter, r *http.Request) { + id, err := strconv.ParseInt(r.PathValue("id"), 10, 64) + if err != nil { + fail(w, 400, "无效的 ID") + return + } + + db.DB.Exec("DELETE FROM tag_options WHERE id=?", id) + success(w, nil) +} diff --git a/internal/llm/client.go b/internal/llm/client.go new file mode 100644 index 0000000..131ba9f --- /dev/null +++ b/internal/llm/client.go @@ -0,0 +1,147 @@ +package llm + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" + + "prompt-generator/internal/models" +) + +type Client struct { + BaseURL string + APIKey string + Model string +} + +func NewClient(baseURL, apiKey, model string) *Client { + return &Client{ + BaseURL: baseURL, + APIKey: apiKey, + Model: model, + } +} + +type chatRequest struct { + Model string `json:"model"` + Messages []chatMessage `json:"messages"` +} + +type chatMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type chatResponse struct { + Choices []struct { + Message struct { + Content string `json:"content"` + } `json:"message"` + } `json:"choices"` + Error *struct { + Message string `json:"message"` + } `json:"error,omitempty"` +} + +func (c *Client) GetSuggestions(currentPrompt string, claudeMD string) ([]models.Suggestion, error) { + systemPrompt := `你是一个 Prompt 工程专家。用户正在构建一个用于 Coding Agent 的 Prompt。 +请分析当前 Prompt,找出缺失的关键约束或可以改进的地方。 +以 JSON 数组返回建议,每条建议包含: +- title: 建议标题 +- description: 详细说明 +- constraint_text: 建议注入的约束文本 + +只返回 JSON 数组,不要有其他内容。` + + userContent := fmt.Sprintf("当前 Prompt 内容如下:\n---\n%s\n---", currentPrompt) + if claudeMD != "" { + userContent += fmt.Sprintf("\n\n用户的 CLAUDE.md 内容:\n---\n%s\n---", claudeMD) + } + + reqBody := chatRequest{ + Model: c.Model, + Messages: []chatMessage{ + {Role: "system", Content: systemPrompt}, + {Role: "user", Content: userContent}, + }, + } + + jsonData, err := json.Marshal(reqBody) + if err != nil { + return nil, fmt.Errorf("序列化请求失败: %w", err) + } + + url := strings.TrimRight(c.BaseURL, "/") + "/chat/completions" + req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) + if err != nil { + return nil, fmt.Errorf("创建请求失败: %w", err) + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+c.APIKey) + + client := &http.Client{Timeout: 30 * time.Second} + resp, err := client.Do(req) + if err != nil { + return nil, fmt.Errorf("调用 LLM API 失败: %w", err) + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("读取响应失败: %w", err) + } + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("LLM API 返回错误 (%d): %s", resp.StatusCode, string(body)) + } + + var chatResp chatResponse + if err := json.Unmarshal(body, &chatResp); err != nil { + return nil, fmt.Errorf("解析响应失败: %w", err) + } + + if chatResp.Error != nil { + return nil, fmt.Errorf("LLM 错误: %s", chatResp.Error.Message) + } + + if len(chatResp.Choices) == 0 { + return nil, fmt.Errorf("LLM 返回空结果") + } + + content := chatResp.Choices[0].Message.Content + suggestions := parseSuggestions(content) + if suggestions == nil { + return nil, fmt.Errorf("AI 返回格式异常,请重试") + } + + return suggestions, nil +} + +func parseSuggestions(data string) []models.Suggestion { + // Try to find JSON array in the response + start := -1 + end := -1 + for i := 0; i < len(data); i++ { + if data[i] == '[' && start == -1 { + start = i + } + if data[i] == ']' { + end = i + 1 + } + } + + if start == -1 || end == -1 { + return nil + } + + var suggestions []models.Suggestion + if err := json.Unmarshal([]byte(data[start:end]), &suggestions); err != nil { + return nil + } + return suggestions +} diff --git a/internal/models/models.go b/internal/models/models.go new file mode 100644 index 0000000..b9db702 --- /dev/null +++ b/internal/models/models.go @@ -0,0 +1,91 @@ +package models + +import "time" + +// Tag 标签定义 +type Tag struct { + ID int64 `json:"id"` + Name string `json:"name"` + Description string `json:"description"` + Scope string `json:"scope"` // system or personal + CreatedAt time.Time `json:"created_at"` + Options []TagOption `json:"options,omitempty"` +} + +// TagOption 标签选项 +type TagOption struct { + ID int64 `json:"id"` + TagID int64 `json:"tag_id"` + Label string `json:"label"` + ConstraintText string `json:"constraint_text"` + SortOrder int `json:"sort_order"` + CreatedAt time.Time `json:"created_at"` +} + +// Snippet 片段模板 +type Snippet struct { + ID int64 `json:"id"` + Name string `json:"name"` + Content string `json:"content"` + Category string `json:"category"` + Scope string `json:"scope"` + CreatedAt time.Time `json:"created_at"` +} + +// BuilderSession 构建会话 +type BuilderSession struct { + ID int64 `json:"id"` + Title string `json:"title"` + ProjectName string `json:"project_name"` + FinalPrompt string `json:"final_prompt"` + ClaudeSessionID *string `json:"claude_session_id"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + Tags []SessionTag `json:"tags,omitempty"` + Snippets []SessionSnippet `json:"snippets,omitempty"` +} + +// SessionTag 构建会话标签关联 +type SessionTag struct { + ID int64 `json:"id"` + BuilderSessionID int64 `json:"builder_session_id"` + TagID int64 `json:"tag_id"` + TagOptionID int64 `json:"tag_option_id"` +} + +// SessionSnippet 构建会话片段关联 +type SessionSnippet struct { + ID int64 `json:"id"` + BuilderSessionID int64 `json:"builder_session_id"` + SnippetID int64 `json:"snippet_id"` + SortOrder int `json:"sort_order"` +} + +// Prompt prompts 表记录(只读) +type Prompt struct { + ID int64 `json:"id"` + SessionID string `json:"session_id"` + ProjectName string `json:"project_name"` + Prompt string `json:"prompt"` + CreatedAt time.Time `json:"created_at"` +} + +// Suggestion 智能建议 +type Suggestion struct { + Title string `json:"title"` + Description string `json:"description"` + ConstraintText string `json:"constraint_text"` +} + +// Settings 系统设置 +type Settings struct { + LLMAPIBaseURL string `json:"llm_api_base_url"` + LLMModelName string `json:"llm_model_name"` +} + +// APIResponse 统一响应 +type APIResponse struct { + Code int `json:"code"` + Message string `json:"message"` + Data interface{} `json:"data,omitempty"` +} diff --git a/main.go b/main.go new file mode 100644 index 0000000..40fed3c --- /dev/null +++ b/main.go @@ -0,0 +1,133 @@ +package main + +import ( + "log" + "net/http" + "os" + + "prompt-generator/internal/auth" + "prompt-generator/internal/config" + "prompt-generator/internal/db" + "prompt-generator/internal/handlers" +) + +func main() { + // Load .env file if exists + loadEnvFile(".env") + + cfg := config.Load() + + // Initialize database + if err := db.Init(cfg); err != nil { + log.Fatalf("Failed to initialize database: %v", err) + } + + // Auto migrate + if err := db.AutoMigrate(); err != nil { + log.Fatalf("Failed to migrate database: %v", err) + } + + // Seed data + if err := db.SeedData(); err != nil { + log.Fatalf("Failed to seed data: %v", err) + } + + // Initialize auth + auth.Init(cfg.SessionSecret) + + // Initialize handlers + handlers.Init(cfg) + + // Setup routes + mux := http.NewServeMux() + + // Static files (frontend) + fs := http.FileServer(http.Dir("frontend")) + mux.Handle("/", fs) + + // Auth routes (no auth required) + mux.HandleFunc("POST /api/auth/login", handlers.Login) + mux.HandleFunc("POST /api/auth/logout", handlers.Logout) + mux.HandleFunc("GET /api/auth/check", handlers.AuthCheck) + + // Protected API routes + apiMux := http.NewServeMux() + apiMux.HandleFunc("GET /api/tags", handlers.GetTags) + apiMux.HandleFunc("POST /api/tags", handlers.CreateTag) + apiMux.HandleFunc("PUT /api/tags/{id}", handlers.UpdateTag) + apiMux.HandleFunc("DELETE /api/tags/{id}", handlers.DeleteTag) + apiMux.HandleFunc("POST /api/tags/{id}/options", handlers.CreateTagOption) + apiMux.HandleFunc("PUT /api/tag-options/{id}", handlers.UpdateTagOption) + apiMux.HandleFunc("DELETE /api/tag-options/{id}", handlers.DeleteTagOption) + + apiMux.HandleFunc("GET /api/snippets", handlers.GetSnippets) + apiMux.HandleFunc("POST /api/snippets", handlers.CreateSnippet) + apiMux.HandleFunc("PUT /api/snippets/{id}", handlers.UpdateSnippet) + apiMux.HandleFunc("DELETE /api/snippets/{id}", handlers.DeleteSnippet) + + apiMux.HandleFunc("GET /api/builder/sessions", handlers.GetBuilderSessions) + apiMux.HandleFunc("POST /api/builder/sessions", handlers.CreateBuilderSession) + apiMux.HandleFunc("GET /api/builder/sessions/{id}", handlers.GetBuilderSession) + apiMux.HandleFunc("PUT /api/builder/sessions/{id}", handlers.UpdateBuilderSession) + apiMux.HandleFunc("DELETE /api/builder/sessions/{id}", handlers.DeleteBuilderSession) + + apiMux.HandleFunc("GET /api/claude/sessions", handlers.GetClaudeSessions) + apiMux.HandleFunc("GET /api/claude/sessions/{session_id}/prompts", handlers.GetClaudePrompts) + + apiMux.HandleFunc("GET /api/dashboard/projects", handlers.GetDashboardProjects) + apiMux.HandleFunc("GET /api/dashboard/projects/{name}/sessions", handlers.GetDashboardSessions) + apiMux.HandleFunc("GET /api/dashboard/prompts", handlers.GetDashboardPrompts) + + apiMux.HandleFunc("POST /api/suggestions", handlers.GetSuggestions) + + apiMux.HandleFunc("GET /api/settings", handlers.GetSettings) + apiMux.HandleFunc("PUT /api/settings", handlers.UpdateSettings) + + // Wrap API routes with auth middleware + mux.Handle("/api/", auth.AuthMiddleware(apiMux)) + + log.Printf("Server starting on :%s", cfg.ServerPort) + if err := http.ListenAndServe(":"+cfg.ServerPort, mux); err != nil { + log.Fatalf("Server failed: %v", err) + } +} + +func loadEnvFile(path string) { + data, err := os.ReadFile(path) + if err != nil { + return // .env file is optional + } + lines := splitLines(string(data)) + for _, line := range lines { + if len(line) == 0 || line[0] == '#' { + continue + } + for i := 0; i < len(line); i++ { + if line[i] == '=' { + key := line[:i] + value := line[i+1:] + // Remove quotes + if len(value) >= 2 && (value[0] == '"' && value[len(value)-1] == '"') { + value = value[1 : len(value)-1] + } + os.Setenv(key, value) + break + } + } + } +} + +func splitLines(s string) []string { + var lines []string + start := 0 + for i := 0; i < len(s); i++ { + if s[i] == '\n' { + lines = append(lines, s[start:i]) + start = i + 1 + } + } + if start < len(s) { + lines = append(lines, s[start:]) + } + return lines +}