feat: add REST API with OpenAPI 3.1 specification
Add webhook/REST API with bearer token auth, chat endpoint, session listing, health check, and embedded OpenAPI spec serving. Includes auth middleware, request validation, and comprehensive tests. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
bbb4266060
commit
3cb176c9fb
9 changed files with 622 additions and 0 deletions
89
api/openapi.yaml
Normal file
89
api/openapi.yaml
Normal file
|
|
@ -0,0 +1,89 @@
|
||||||
|
openapi: "3.1.0"
|
||||||
|
info:
|
||||||
|
title: PicoClaw API
|
||||||
|
version: "1.0.0"
|
||||||
|
description: REST API for PicoClaw AI Gateway
|
||||||
|
servers:
|
||||||
|
- url: http://localhost:18790
|
||||||
|
security:
|
||||||
|
- ApiKeyAuth: []
|
||||||
|
components:
|
||||||
|
securitySchemes:
|
||||||
|
ApiKeyAuth:
|
||||||
|
type: http
|
||||||
|
scheme: bearer
|
||||||
|
schemas:
|
||||||
|
ChatRequest:
|
||||||
|
type: object
|
||||||
|
required: [message]
|
||||||
|
properties:
|
||||||
|
session_id: { type: string, description: "Session ID (auto-generated if omitted)" }
|
||||||
|
message: { type: string, description: "User message" }
|
||||||
|
ChatResponse:
|
||||||
|
type: object
|
||||||
|
properties:
|
||||||
|
session_id: { type: string }
|
||||||
|
response: { type: string }
|
||||||
|
SessionList:
|
||||||
|
type: object
|
||||||
|
properties:
|
||||||
|
sessions: { type: array, items: { type: string } }
|
||||||
|
HealthResponse:
|
||||||
|
type: object
|
||||||
|
properties:
|
||||||
|
status: { type: string, enum: [ok] }
|
||||||
|
version: { type: string }
|
||||||
|
Error:
|
||||||
|
type: object
|
||||||
|
properties:
|
||||||
|
error: { type: string }
|
||||||
|
code: { type: integer }
|
||||||
|
paths:
|
||||||
|
/api/chat:
|
||||||
|
post:
|
||||||
|
operationId: sendMessage
|
||||||
|
summary: Send a message and get an AI response
|
||||||
|
requestBody:
|
||||||
|
required: true
|
||||||
|
content:
|
||||||
|
application/json:
|
||||||
|
schema: { $ref: "#/components/schemas/ChatRequest" }
|
||||||
|
responses:
|
||||||
|
"200":
|
||||||
|
description: Successful response
|
||||||
|
content:
|
||||||
|
application/json:
|
||||||
|
schema: { $ref: "#/components/schemas/ChatResponse" }
|
||||||
|
"401": { description: Unauthorized }
|
||||||
|
"500": { description: Internal error }
|
||||||
|
/api/sessions:
|
||||||
|
get:
|
||||||
|
operationId: listSessions
|
||||||
|
summary: List active sessions
|
||||||
|
responses:
|
||||||
|
"200":
|
||||||
|
description: Successful response
|
||||||
|
content:
|
||||||
|
application/json:
|
||||||
|
schema: { $ref: "#/components/schemas/SessionList" }
|
||||||
|
/api/health:
|
||||||
|
get:
|
||||||
|
operationId: healthCheck
|
||||||
|
summary: Health check
|
||||||
|
security: []
|
||||||
|
responses:
|
||||||
|
"200":
|
||||||
|
description: Successful response
|
||||||
|
content:
|
||||||
|
application/json:
|
||||||
|
schema: { $ref: "#/components/schemas/HealthResponse" }
|
||||||
|
/api/openapi.yaml:
|
||||||
|
get:
|
||||||
|
operationId: getSpec
|
||||||
|
summary: Serve the OpenAPI specification
|
||||||
|
security: []
|
||||||
|
responses:
|
||||||
|
"200":
|
||||||
|
description: OpenAPI spec
|
||||||
|
content:
|
||||||
|
application/x-yaml: {}
|
||||||
89
cmd/picoclaw/api/openapi.yaml
Normal file
89
cmd/picoclaw/api/openapi.yaml
Normal file
|
|
@ -0,0 +1,89 @@
|
||||||
|
openapi: "3.1.0"
|
||||||
|
info:
|
||||||
|
title: PicoClaw API
|
||||||
|
version: "1.0.0"
|
||||||
|
description: REST API for PicoClaw AI Gateway
|
||||||
|
servers:
|
||||||
|
- url: http://localhost:18790
|
||||||
|
security:
|
||||||
|
- ApiKeyAuth: []
|
||||||
|
components:
|
||||||
|
securitySchemes:
|
||||||
|
ApiKeyAuth:
|
||||||
|
type: http
|
||||||
|
scheme: bearer
|
||||||
|
schemas:
|
||||||
|
ChatRequest:
|
||||||
|
type: object
|
||||||
|
required: [message]
|
||||||
|
properties:
|
||||||
|
session_id: { type: string, description: "Session ID (auto-generated if omitted)" }
|
||||||
|
message: { type: string, description: "User message" }
|
||||||
|
ChatResponse:
|
||||||
|
type: object
|
||||||
|
properties:
|
||||||
|
session_id: { type: string }
|
||||||
|
response: { type: string }
|
||||||
|
SessionList:
|
||||||
|
type: object
|
||||||
|
properties:
|
||||||
|
sessions: { type: array, items: { type: string } }
|
||||||
|
HealthResponse:
|
||||||
|
type: object
|
||||||
|
properties:
|
||||||
|
status: { type: string, enum: [ok] }
|
||||||
|
version: { type: string }
|
||||||
|
Error:
|
||||||
|
type: object
|
||||||
|
properties:
|
||||||
|
error: { type: string }
|
||||||
|
code: { type: integer }
|
||||||
|
paths:
|
||||||
|
/api/chat:
|
||||||
|
post:
|
||||||
|
operationId: sendMessage
|
||||||
|
summary: Send a message and get an AI response
|
||||||
|
requestBody:
|
||||||
|
required: true
|
||||||
|
content:
|
||||||
|
application/json:
|
||||||
|
schema: { $ref: "#/components/schemas/ChatRequest" }
|
||||||
|
responses:
|
||||||
|
"200":
|
||||||
|
description: Successful response
|
||||||
|
content:
|
||||||
|
application/json:
|
||||||
|
schema: { $ref: "#/components/schemas/ChatResponse" }
|
||||||
|
"401": { description: Unauthorized }
|
||||||
|
"500": { description: Internal error }
|
||||||
|
/api/sessions:
|
||||||
|
get:
|
||||||
|
operationId: listSessions
|
||||||
|
summary: List active sessions
|
||||||
|
responses:
|
||||||
|
"200":
|
||||||
|
description: Successful response
|
||||||
|
content:
|
||||||
|
application/json:
|
||||||
|
schema: { $ref: "#/components/schemas/SessionList" }
|
||||||
|
/api/health:
|
||||||
|
get:
|
||||||
|
operationId: healthCheck
|
||||||
|
summary: Health check
|
||||||
|
security: []
|
||||||
|
responses:
|
||||||
|
"200":
|
||||||
|
description: Successful response
|
||||||
|
content:
|
||||||
|
application/json:
|
||||||
|
schema: { $ref: "#/components/schemas/HealthResponse" }
|
||||||
|
/api/openapi.yaml:
|
||||||
|
get:
|
||||||
|
operationId: getSpec
|
||||||
|
summary: Serve the OpenAPI specification
|
||||||
|
security: []
|
||||||
|
responses:
|
||||||
|
"200":
|
||||||
|
description: OpenAPI spec
|
||||||
|
content:
|
||||||
|
application/x-yaml: {}
|
||||||
|
|
@ -13,6 +13,7 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"io/fs"
|
"io/fs"
|
||||||
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"os/signal"
|
"os/signal"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
|
@ -28,6 +29,7 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
"github.com/sipeed/picoclaw/pkg/cron"
|
"github.com/sipeed/picoclaw/pkg/cron"
|
||||||
"github.com/sipeed/picoclaw/pkg/devices"
|
"github.com/sipeed/picoclaw/pkg/devices"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/gateway"
|
||||||
"github.com/sipeed/picoclaw/pkg/heartbeat"
|
"github.com/sipeed/picoclaw/pkg/heartbeat"
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
"github.com/sipeed/picoclaw/pkg/migrate"
|
"github.com/sipeed/picoclaw/pkg/migrate"
|
||||||
|
|
@ -43,6 +45,10 @@ import (
|
||||||
//go:embed workspace
|
//go:embed workspace
|
||||||
var embeddedFiles embed.FS
|
var embeddedFiles embed.FS
|
||||||
|
|
||||||
|
//go:generate cp -r ../../api .
|
||||||
|
//go:embed api/openapi.yaml
|
||||||
|
var openapiSpec []byte
|
||||||
|
|
||||||
var (
|
var (
|
||||||
version = "dev"
|
version = "dev"
|
||||||
gitCommit string
|
gitCommit string
|
||||||
|
|
@ -680,6 +686,30 @@ func gatewayCmd() {
|
||||||
fmt.Printf("Error starting channels: %v\n", err)
|
fmt.Printf("Error starting channels: %v\n", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Start REST API if enabled
|
||||||
|
if cfg.API.Enabled {
|
||||||
|
gateway.SetVersion(version)
|
||||||
|
gateway.SetOpenAPISpec(openapiSpec)
|
||||||
|
apiServer := gateway.NewAPIServer(agentLoop, cfg.API)
|
||||||
|
apiAddr := fmt.Sprintf("%s:%d", cfg.Gateway.Host, cfg.Gateway.Port)
|
||||||
|
httpServer := &http.Server{
|
||||||
|
Addr: apiAddr,
|
||||||
|
Handler: apiServer.Handler(),
|
||||||
|
}
|
||||||
|
go func() {
|
||||||
|
fmt.Printf("✓ REST API listening on %s\n", apiAddr)
|
||||||
|
if err := httpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
||||||
|
logger.ErrorCF("api", "REST API server error", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
go func() {
|
||||||
|
<-ctx.Done()
|
||||||
|
httpServer.Shutdown(context.Background())
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
go agentLoop.Run(ctx)
|
go agentLoop.Run(ctx)
|
||||||
|
|
||||||
sigChan := make(chan os.Signal, 1)
|
sigChan := make(chan os.Signal, 1)
|
||||||
|
|
|
||||||
|
|
@ -203,6 +203,11 @@ func (al *AgentLoop) RegisterTool(tool tools.Tool) {
|
||||||
al.tools.Register(tool)
|
al.tools.Register(tool)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetSessionManager returns the session manager for external access (e.g. REST API).
|
||||||
|
func (al *AgentLoop) GetSessionManager() *session.SessionManager {
|
||||||
|
return al.sessions
|
||||||
|
}
|
||||||
|
|
||||||
// RecordLastChannel records the last active channel for this workspace.
|
// RecordLastChannel records the last active channel for this workspace.
|
||||||
// This uses the atomic state save mechanism to prevent data loss on crash.
|
// This uses the atomic state save mechanism to prevent data loss on crash.
|
||||||
func (al *AgentLoop) RecordLastChannel(channel string) error {
|
func (al *AgentLoop) RecordLastChannel(channel string) error {
|
||||||
|
|
|
||||||
|
|
@ -52,6 +52,7 @@ type Config struct {
|
||||||
Heartbeat HeartbeatConfig `json:"heartbeat"`
|
Heartbeat HeartbeatConfig `json:"heartbeat"`
|
||||||
Tracing TracingConfig `json:"tracing"`
|
Tracing TracingConfig `json:"tracing"`
|
||||||
Devices DevicesConfig `json:"devices"`
|
Devices DevicesConfig `json:"devices"`
|
||||||
|
API APIConfig `json:"api"`
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -201,6 +202,11 @@ type GatewayConfig struct {
|
||||||
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type APIConfig struct {
|
||||||
|
Enabled bool `json:"enabled" env:"PICOCLAW_API_ENABLED"`
|
||||||
|
APIKey string `json:"api_key" env:"PICOCLAW_API_KEY"`
|
||||||
|
}
|
||||||
|
|
||||||
type BraveConfig struct {
|
type BraveConfig struct {
|
||||||
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_BRAVE_ENABLED"`
|
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_BRAVE_ENABLED"`
|
||||||
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEY"`
|
APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEY"`
|
||||||
|
|
@ -341,6 +347,10 @@ func DefaultConfig() *Config {
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
MonitorUSB: true,
|
MonitorUSB: true,
|
||||||
},
|
},
|
||||||
|
API: APIConfig{
|
||||||
|
Enabled: false,
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
166
pkg/gateway/api.go
Normal file
166
pkg/gateway/api.go
Normal file
|
|
@ -0,0 +1,166 @@
|
||||||
|
package gateway
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/agent"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// openapiSpec holds the OpenAPI YAML spec, set via SetOpenAPISpec.
|
||||||
|
var openapiSpec []byte
|
||||||
|
|
||||||
|
// SetOpenAPISpec sets the embedded OpenAPI spec content.
|
||||||
|
func SetOpenAPISpec(spec []byte) {
|
||||||
|
openapiSpec = spec
|
||||||
|
}
|
||||||
|
|
||||||
|
// version is set by the caller (main.go) via SetVersion.
|
||||||
|
var apiVersion = "dev"
|
||||||
|
|
||||||
|
// SetVersion sets the version string returned by the health endpoint.
|
||||||
|
func SetVersion(v string) {
|
||||||
|
apiVersion = v
|
||||||
|
}
|
||||||
|
|
||||||
|
// API request/response types matching OpenAPI schemas
|
||||||
|
type ChatRequest struct {
|
||||||
|
SessionID string `json:"session_id,omitempty"`
|
||||||
|
Message string `json:"message"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type ChatResponse struct {
|
||||||
|
SessionID string `json:"session_id"`
|
||||||
|
Response string `json:"response"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type SessionList struct {
|
||||||
|
Sessions []string `json:"sessions"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type HealthResponse struct {
|
||||||
|
Status string `json:"status"`
|
||||||
|
Version string `json:"version"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type ErrorResponse struct {
|
||||||
|
Error string `json:"error"`
|
||||||
|
Code int `json:"code"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// APIServer holds the dependencies for API handlers.
|
||||||
|
type APIServer struct {
|
||||||
|
agentLoop *agent.AgentLoop
|
||||||
|
config config.APIConfig
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewAPIServer creates a new API server.
|
||||||
|
func NewAPIServer(agentLoop *agent.AgentLoop, cfg config.APIConfig) *APIServer {
|
||||||
|
return &APIServer{
|
||||||
|
agentLoop: agentLoop,
|
||||||
|
config: cfg,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handler returns an http.Handler with all API routes and auth middleware.
|
||||||
|
func (s *APIServer) Handler() http.Handler {
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
|
||||||
|
mux.HandleFunc("/api/chat", s.handleChat)
|
||||||
|
mux.HandleFunc("/api/sessions", s.handleSessions)
|
||||||
|
mux.HandleFunc("/api/health", s.handleHealth)
|
||||||
|
mux.HandleFunc("/api/openapi.yaml", s.handleSpec)
|
||||||
|
|
||||||
|
publicPaths := []string{"/api/health", "/api/openapi.yaml"}
|
||||||
|
return AuthMiddleware(s.config.APIKey, publicPaths, mux)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *APIServer) handleChat(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != http.MethodPost {
|
||||||
|
writeError(w, http.StatusMethodNotAllowed, "method not allowed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var req ChatRequest
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||||
|
writeError(w, http.StatusBadRequest, "invalid request body")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if req.Message == "" {
|
||||||
|
writeError(w, http.StatusBadRequest, "message is required")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionID := req.SessionID
|
||||||
|
if sessionID == "" {
|
||||||
|
sessionID = fmt.Sprintf("api:%s", uuid.New().String())
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.InfoCF("api", "Chat request", map[string]interface{}{
|
||||||
|
"session_id": sessionID,
|
||||||
|
"message": req.Message,
|
||||||
|
})
|
||||||
|
|
||||||
|
response, err := s.agentLoop.ProcessDirect(context.Background(), req.Message, sessionID)
|
||||||
|
if err != nil {
|
||||||
|
logger.ErrorCF("api", "Chat error", map[string]interface{}{
|
||||||
|
"error": err.Error(),
|
||||||
|
})
|
||||||
|
writeError(w, http.StatusInternalServerError, "failed to process message")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
writeJSON(w, http.StatusOK, ChatResponse{
|
||||||
|
SessionID: sessionID,
|
||||||
|
Response: response,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *APIServer) handleSessions(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != http.MethodGet {
|
||||||
|
writeError(w, http.StatusMethodNotAllowed, "method not allowed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
sessions := s.agentLoop.GetSessionManager().ListSessions()
|
||||||
|
writeJSON(w, http.StatusOK, SessionList{Sessions: sessions})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *APIServer) handleHealth(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != http.MethodGet {
|
||||||
|
writeError(w, http.StatusMethodNotAllowed, "method not allowed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
writeJSON(w, http.StatusOK, HealthResponse{
|
||||||
|
Status: "ok",
|
||||||
|
Version: apiVersion,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *APIServer) handleSpec(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != http.MethodGet {
|
||||||
|
writeError(w, http.StatusMethodNotAllowed, "method not allowed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
w.Header().Set("Content-Type", "application/x-yaml")
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
w.Write(openapiSpec)
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeJSON(w http.ResponseWriter, status int, v interface{}) {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(status)
|
||||||
|
json.NewEncoder(w).Encode(v)
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeError(w http.ResponseWriter, code int, message string) {
|
||||||
|
writeJSON(w, code, ErrorResponse{Error: message, Code: code})
|
||||||
|
}
|
||||||
175
pkg/gateway/api_test.go
Normal file
175
pkg/gateway/api_test.go
Normal file
|
|
@ -0,0 +1,175 @@
|
||||||
|
package gateway
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAuthMiddleware(t *testing.T) {
|
||||||
|
inner := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
w.Write([]byte("ok"))
|
||||||
|
})
|
||||||
|
|
||||||
|
publicPaths := []string{"/api/health"}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
apiKey string
|
||||||
|
authHeader string
|
||||||
|
path string
|
||||||
|
wantStatus int
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "valid bearer token",
|
||||||
|
apiKey: "secret123",
|
||||||
|
authHeader: "Bearer secret123",
|
||||||
|
path: "/api/chat",
|
||||||
|
wantStatus: http.StatusOK,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "invalid bearer token",
|
||||||
|
apiKey: "secret123",
|
||||||
|
authHeader: "Bearer wrong",
|
||||||
|
path: "/api/chat",
|
||||||
|
wantStatus: http.StatusUnauthorized,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "missing auth header",
|
||||||
|
apiKey: "secret123",
|
||||||
|
authHeader: "",
|
||||||
|
path: "/api/chat",
|
||||||
|
wantStatus: http.StatusUnauthorized,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "public path skips auth",
|
||||||
|
apiKey: "secret123",
|
||||||
|
authHeader: "",
|
||||||
|
path: "/api/health",
|
||||||
|
wantStatus: http.StatusOK,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no api key configured (open mode)",
|
||||||
|
apiKey: "",
|
||||||
|
authHeader: "",
|
||||||
|
path: "/api/chat",
|
||||||
|
wantStatus: http.StatusOK,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "wrong auth format",
|
||||||
|
apiKey: "secret123",
|
||||||
|
authHeader: "Basic dXNlcjpwYXNz",
|
||||||
|
path: "/api/chat",
|
||||||
|
wantStatus: http.StatusUnauthorized,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
handler := AuthMiddleware(tt.apiKey, publicPaths, inner)
|
||||||
|
req := httptest.NewRequest(http.MethodGet, tt.path, nil)
|
||||||
|
if tt.authHeader != "" {
|
||||||
|
req.Header.Set("Authorization", tt.authHeader)
|
||||||
|
}
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
handler.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
if w.Code != tt.wantStatus {
|
||||||
|
t.Errorf("status = %d, want %d", w.Code, tt.wantStatus)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHealthEndpoint(t *testing.T) {
|
||||||
|
SetVersion("1.2.3")
|
||||||
|
s := &APIServer{}
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/health", nil)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
s.handleHealth(w, req)
|
||||||
|
|
||||||
|
resp := w.Result()
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusOK)
|
||||||
|
}
|
||||||
|
|
||||||
|
var health HealthResponse
|
||||||
|
if err := json.NewDecoder(resp.Body).Decode(&health); err != nil {
|
||||||
|
t.Fatalf("decode: %v", err)
|
||||||
|
}
|
||||||
|
if health.Status != "ok" {
|
||||||
|
t.Errorf("status = %q, want %q", health.Status, "ok")
|
||||||
|
}
|
||||||
|
if health.Version != "1.2.3" {
|
||||||
|
t.Errorf("version = %q, want %q", health.Version, "1.2.3")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSpecEndpoint(t *testing.T) {
|
||||||
|
SetOpenAPISpec([]byte("openapi: '3.1.0'\ninfo:\n title: test"))
|
||||||
|
s := &APIServer{}
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/openapi.yaml", nil)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
s.handleSpec(w, req)
|
||||||
|
|
||||||
|
resp := w.Result()
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusOK)
|
||||||
|
}
|
||||||
|
|
||||||
|
ct := resp.Header.Get("Content-Type")
|
||||||
|
if ct != "application/x-yaml" {
|
||||||
|
t.Errorf("content-type = %q, want %q", ct, "application/x-yaml")
|
||||||
|
}
|
||||||
|
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
if !strings.Contains(string(body), "openapi") {
|
||||||
|
t.Error("response body should contain 'openapi'")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestChatEndpointRequiresMessage(t *testing.T) {
|
||||||
|
s := &APIServer{}
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/chat", strings.NewReader(`{}`))
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
s.handleChat(w, req)
|
||||||
|
|
||||||
|
resp := w.Result()
|
||||||
|
if resp.StatusCode != http.StatusBadRequest {
|
||||||
|
t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusBadRequest)
|
||||||
|
}
|
||||||
|
|
||||||
|
var errResp ErrorResponse
|
||||||
|
json.NewDecoder(resp.Body).Decode(&errResp)
|
||||||
|
if errResp.Error != "message is required" {
|
||||||
|
t.Errorf("error = %q, want %q", errResp.Error, "message is required")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestChatEndpointRejectsGet(t *testing.T) {
|
||||||
|
s := &APIServer{}
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/chat", nil)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
s.handleChat(w, req)
|
||||||
|
|
||||||
|
resp := w.Result()
|
||||||
|
if resp.StatusCode != http.StatusMethodNotAllowed {
|
||||||
|
t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusMethodNotAllowed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestChatEndpointInvalidJSON(t *testing.T) {
|
||||||
|
s := &APIServer{}
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/chat", strings.NewReader("not json"))
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
s.handleChat(w, req)
|
||||||
|
|
||||||
|
resp := w.Result()
|
||||||
|
if resp.StatusCode != http.StatusBadRequest {
|
||||||
|
t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusBadRequest)
|
||||||
|
}
|
||||||
|
}
|
||||||
46
pkg/gateway/middleware.go
Normal file
46
pkg/gateway/middleware.go
Normal file
|
|
@ -0,0 +1,46 @@
|
||||||
|
package gateway
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AuthMiddleware validates the Authorization: Bearer <token> header
|
||||||
|
// against the configured API key. Public paths bypass auth.
|
||||||
|
func AuthMiddleware(apiKey string, publicPaths []string, next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Skip auth for public paths
|
||||||
|
for _, p := range publicPaths {
|
||||||
|
if r.URL.Path == p {
|
||||||
|
next.ServeHTTP(w, r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Skip auth if no API key configured (open mode)
|
||||||
|
if apiKey == "" {
|
||||||
|
next.ServeHTTP(w, r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
auth := r.Header.Get("Authorization")
|
||||||
|
if auth == "" {
|
||||||
|
writeError(w, http.StatusUnauthorized, "missing authorization header")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
const prefix = "Bearer "
|
||||||
|
if !strings.HasPrefix(auth, prefix) {
|
||||||
|
writeError(w, http.StatusUnauthorized, "invalid authorization format")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
token := auth[len(prefix):]
|
||||||
|
if token != apiKey {
|
||||||
|
writeError(w, http.StatusUnauthorized, "invalid API key")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
next.ServeHTTP(w, r)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
@ -86,6 +86,18 @@ func (sm *SessionManager) AddFullMessage(sessionKey string, msg providers.Messag
|
||||||
session.Updated = time.Now()
|
session.Updated = time.Now()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ListSessions returns all session keys.
|
||||||
|
func (sm *SessionManager) ListSessions() []string {
|
||||||
|
sm.mu.RLock()
|
||||||
|
defer sm.mu.RUnlock()
|
||||||
|
|
||||||
|
keys := make([]string, 0, len(sm.sessions))
|
||||||
|
for k := range sm.sessions {
|
||||||
|
keys = append(keys, k)
|
||||||
|
}
|
||||||
|
return keys
|
||||||
|
}
|
||||||
|
|
||||||
func (sm *SessionManager) GetHistory(key string) []providers.Message {
|
func (sm *SessionManager) GetHistory(key string) []providers.Message {
|
||||||
sm.mu.RLock()
|
sm.mu.RLock()
|
||||||
defer sm.mu.RUnlock()
|
defer sm.mu.RUnlock()
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue