feat: implement attachments with content-addressable storage
Add file attachment support with SHA-256 content-addressable storage, automatic deduplication, MIME detection, and garbage collection for orphaned files. Includes MCP tools (upload_attachment, download_attachment, gc_attachments), REST API endpoints for Web UI, and comprehensive tests. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
1cb0bb7a8f
commit
8e2294e19d
+15
-3
@@ -9,6 +9,7 @@ import (
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
@@ -19,6 +20,7 @@ import (
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/agents"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/api"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/attachments"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/auth"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/channels"
|
||||
mcpserver "github.com/smart-mcp-proxy/synapbus/internal/mcp"
|
||||
@@ -180,6 +182,16 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
channelStore := channels.NewSQLiteChannelStore(db.DB)
|
||||
channelService := channels.NewService(channelStore, msgService, tracer)
|
||||
|
||||
// Create attachment service
|
||||
attachmentsDir := filepath.Join(dataDir, "attachments")
|
||||
cas, err := attachments.NewCAS(attachmentsDir, slog.Default())
|
||||
if err != nil {
|
||||
return fmt.Errorf("create CAS engine: %w", err)
|
||||
}
|
||||
attachmentStore := attachments.NewSQLiteStore(db.DB, slog.Default())
|
||||
attachmentService := attachments.NewService(attachmentStore, cas, slog.Default())
|
||||
slog.Info("attachment service initialized", "dir", attachmentsDir)
|
||||
|
||||
// Initialize auth subsystem
|
||||
authSecret := make([]byte, 32)
|
||||
if _, err := rand.Read(authSecret); err != nil {
|
||||
@@ -221,7 +233,7 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
}
|
||||
|
||||
// Create MCP server
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService)
|
||||
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, attachmentService)
|
||||
startTime := time.Now()
|
||||
|
||||
// Set up chi router
|
||||
@@ -250,8 +262,8 @@ func runServe(cmd *cobra.Command, args []string) error {
|
||||
// MCP SSE endpoint
|
||||
r.Mount("/mcp", mcpSrv.SSEHandler())
|
||||
|
||||
// Mount API routes (traces, export, stats, metrics)
|
||||
apiRouter := api.NewRouter(traceStore, metrics)
|
||||
// Mount API routes (traces, export, stats, metrics, attachments)
|
||||
apiRouter := api.NewRouter(traceStore, metrics, attachmentService)
|
||||
r.Mount("/", apiRouter)
|
||||
|
||||
// Start HTTP server
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/attachments"
|
||||
)
|
||||
|
||||
// AttachmentsHandler provides REST API endpoints for attachment operations.
|
||||
// These endpoints are intended for the Web UI, not for agent-to-agent use
|
||||
// (agents use MCP tools instead).
|
||||
type AttachmentsHandler struct {
|
||||
service *attachments.Service
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewAttachmentsHandler creates a new attachments REST handler.
|
||||
func NewAttachmentsHandler(service *attachments.Service) *AttachmentsHandler {
|
||||
return &AttachmentsHandler{
|
||||
service: service,
|
||||
logger: slog.Default().With("component", "api-attachments"),
|
||||
}
|
||||
}
|
||||
|
||||
// Download streams an attachment file to the client.
|
||||
// GET /api/attachments/{hash}
|
||||
func (h *AttachmentsHandler) Download(w http.ResponseWriter, r *http.Request) {
|
||||
hash := chi.URLParam(r, "hash")
|
||||
if hash == "" {
|
||||
http.Error(w, `{"error":"hash parameter required"}`, http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.service.Download(r.Context(), hash)
|
||||
if err != nil {
|
||||
switch err {
|
||||
case attachments.ErrNotFound, attachments.ErrFileMissing:
|
||||
http.Error(w, `{"error":"attachment not found"}`, http.StatusNotFound)
|
||||
default:
|
||||
h.logger.Error("download attachment failed", "hash", hash, "error", err)
|
||||
http.Error(w, `{"error":"internal server error"}`, http.StatusInternalServerError)
|
||||
}
|
||||
return
|
||||
}
|
||||
defer result.Content.Close()
|
||||
|
||||
w.Header().Set("Content-Type", result.MIMEType)
|
||||
if result.Size > 0 {
|
||||
w.Header().Set("Content-Length", fmt.Sprintf("%d", result.Size))
|
||||
}
|
||||
|
||||
// Images are displayed inline; everything else triggers a download.
|
||||
if attachments.IsImageType(result.MIMEType) {
|
||||
w.Header().Set("Content-Disposition", fmt.Sprintf("inline; filename=%q", result.Filename))
|
||||
} else {
|
||||
w.Header().Set("Content-Disposition", fmt.Sprintf("attachment; filename=%q", result.Filename))
|
||||
}
|
||||
|
||||
if _, err := streamContent(w, result.Content); err != nil {
|
||||
h.logger.Error("stream attachment failed", "hash", hash, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Metadata returns attachment metadata as JSON.
|
||||
// GET /api/attachments/{hash}/meta
|
||||
func (h *AttachmentsHandler) Metadata(w http.ResponseWriter, r *http.Request) {
|
||||
hash := chi.URLParam(r, "hash")
|
||||
if hash == "" {
|
||||
http.Error(w, `{"error":"hash parameter required"}`, http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.service.Download(r.Context(), hash)
|
||||
if err != nil {
|
||||
switch err {
|
||||
case attachments.ErrNotFound, attachments.ErrFileMissing:
|
||||
http.Error(w, `{"error":"attachment not found"}`, http.StatusNotFound)
|
||||
default:
|
||||
h.logger.Error("get attachment metadata failed", "hash", hash, "error", err)
|
||||
http.Error(w, `{"error":"internal server error"}`, http.StatusInternalServerError)
|
||||
}
|
||||
return
|
||||
}
|
||||
// Close the content reader immediately since we only need metadata.
|
||||
result.Content.Close()
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"hash": result.Hash,
|
||||
"original_filename": result.Filename,
|
||||
"mime_type": result.MIMEType,
|
||||
"size": result.Size,
|
||||
"is_image": attachments.IsImageType(result.MIMEType),
|
||||
})
|
||||
}
|
||||
|
||||
// Upload handles multipart file uploads from the Web UI.
|
||||
// POST /api/attachments
|
||||
func (h *AttachmentsHandler) Upload(w http.ResponseWriter, r *http.Request) {
|
||||
// Limit request body to MaxFileSize + overhead for multipart headers.
|
||||
r.Body = http.MaxBytesReader(w, r.Body, attachments.MaxFileSize+1024*1024)
|
||||
|
||||
if err := r.ParseMultipartForm(attachments.MaxFileSize); err != nil {
|
||||
http.Error(w, `{"error":"file too large or invalid multipart form"}`, http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
file, header, err := r.FormFile("file")
|
||||
if err != nil {
|
||||
http.Error(w, `{"error":"file field required"}`, http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
// Extract uploader identity from context (set by auth middleware).
|
||||
uploadedBy := "web-ui"
|
||||
if ownerID, ok := OwnerIDFromContext(r.Context()); ok {
|
||||
uploadedBy = fmt.Sprintf("owner-%d", ownerID)
|
||||
}
|
||||
|
||||
req := attachments.UploadRequest{
|
||||
Content: file,
|
||||
Filename: header.Filename,
|
||||
UploadedBy: uploadedBy,
|
||||
}
|
||||
|
||||
result, err := h.service.Upload(r.Context(), req)
|
||||
if err != nil {
|
||||
switch err {
|
||||
case attachments.ErrEmptyFile:
|
||||
http.Error(w, `{"error":"empty file not allowed"}`, http.StatusBadRequest)
|
||||
case attachments.ErrFileTooLarge:
|
||||
http.Error(w, `{"error":"file exceeds maximum size of 50MB"}`, http.StatusRequestEntityTooLarge)
|
||||
default:
|
||||
h.logger.Error("upload attachment failed", "error", err)
|
||||
http.Error(w, `{"error":"internal server error"}`, http.StatusInternalServerError)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"hash": result.Hash,
|
||||
"size": result.Size,
|
||||
"mime_type": result.MIMEType,
|
||||
"original_filename": result.Filename,
|
||||
})
|
||||
}
|
||||
|
||||
// streamContent copies the reader to the response writer.
|
||||
func streamContent(dst http.ResponseWriter, src io.Reader) (int64, error) {
|
||||
return io.Copy(dst, src)
|
||||
}
|
||||
+14
-1
@@ -5,12 +5,14 @@ import (
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/attachments"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/trace"
|
||||
)
|
||||
|
||||
// NewRouter creates a chi router with all API routes configured.
|
||||
// metricsInstance may be nil if metrics are disabled.
|
||||
func NewRouter(traceStore trace.TraceStore, metricsInstance *trace.Metrics) chi.Router {
|
||||
// attachmentService may be nil if attachments are not configured.
|
||||
func NewRouter(traceStore trace.TraceStore, metricsInstance *trace.Metrics, attachmentService *attachments.Service) chi.Router {
|
||||
r := chi.NewRouter()
|
||||
|
||||
// Global middleware
|
||||
@@ -28,6 +30,17 @@ func NewRouter(traceStore trace.TraceStore, metricsInstance *trace.Metrics) chi.
|
||||
r.Get("/api/traces/stats", tracesHandler.TraceStats)
|
||||
})
|
||||
|
||||
// Attachment API routes (for Web UI)
|
||||
if attachmentService != nil {
|
||||
attachmentsHandler := NewAttachmentsHandler(attachmentService)
|
||||
r.Get("/api/attachments/{hash}", attachmentsHandler.Download)
|
||||
r.Get("/api/attachments/{hash}/meta", attachmentsHandler.Metadata)
|
||||
r.Group(func(r chi.Router) {
|
||||
r.Use(OwnerAuthMiddleware)
|
||||
r.Post("/api/attachments", attachmentsHandler.Upload)
|
||||
})
|
||||
}
|
||||
|
||||
// Metrics endpoint (unauthenticated, only registered when enabled)
|
||||
if metricsInstance != nil {
|
||||
r.Get("/metrics", func(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
@@ -81,7 +81,7 @@ func TestListTraces_OwnerIsolation(t *testing.T) {
|
||||
seedTraces(t, store, "1", "alice-bot", []string{"send_message", "read_inbox", "send_message"})
|
||||
seedTraces(t, store, "2", "bob-bot", []string{"send_message", "error"})
|
||||
|
||||
router := NewRouter(store, nil)
|
||||
router := NewRouter(store, nil, nil)
|
||||
|
||||
t.Run("owner 1 sees only own traces", func(t *testing.T) {
|
||||
rr := makeRequest(t, router, "GET", "/api/traces", "1")
|
||||
@@ -141,7 +141,7 @@ func TestListTraces_Filters(t *testing.T) {
|
||||
seedTraces(t, store, "1", "agent-a", []string{"send_message", "read_inbox", "send_message", "error"})
|
||||
seedTraces(t, store, "1", "agent-b", []string{"send_message", "join_channel"})
|
||||
|
||||
router := NewRouter(store, nil)
|
||||
router := NewRouter(store, nil, nil)
|
||||
|
||||
t.Run("filter by agent_name", func(t *testing.T) {
|
||||
rr := makeRequest(t, router, "GET", "/api/traces?agent_name=agent-a", "1")
|
||||
@@ -223,7 +223,7 @@ func TestListTraces_ResponseFormat(t *testing.T) {
|
||||
|
||||
seedTraces(t, store, "1", "agent", []string{"send_message"})
|
||||
|
||||
router := NewRouter(store, nil)
|
||||
router := NewRouter(store, nil, nil)
|
||||
rr := makeRequest(t, router, "GET", "/api/traces", "1")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d", rr.Code)
|
||||
@@ -267,7 +267,7 @@ func TestTraceStats(t *testing.T) {
|
||||
|
||||
seedTraces(t, store, "1", "agent", []string{"send_message", "send_message", "read_inbox", "error"})
|
||||
|
||||
router := NewRouter(store, nil)
|
||||
router := NewRouter(store, nil, nil)
|
||||
rr := makeRequest(t, router, "GET", "/api/traces/stats", "1")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d", rr.Code)
|
||||
@@ -288,7 +288,7 @@ func TestExportTraces_JSON(t *testing.T) {
|
||||
|
||||
seedTraces(t, store, "1", "agent", []string{"send_message", "read_inbox"})
|
||||
|
||||
router := NewRouter(store, nil)
|
||||
router := NewRouter(store, nil, nil)
|
||||
rr := makeRequest(t, router, "GET", "/api/traces/export?format=json", "1")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d", rr.Code)
|
||||
@@ -318,7 +318,7 @@ func TestExportTraces_CSV(t *testing.T) {
|
||||
|
||||
seedTraces(t, store, "1", "agent", []string{"send_message", "read_inbox"})
|
||||
|
||||
router := NewRouter(store, nil)
|
||||
router := NewRouter(store, nil, nil)
|
||||
rr := makeRequest(t, router, "GET", "/api/traces/export?format=csv", "1")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d", rr.Code)
|
||||
@@ -353,7 +353,7 @@ func TestExportTraces_Empty(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
store := trace.NewSQLiteTraceStore(db)
|
||||
|
||||
router := NewRouter(store, nil)
|
||||
router := NewRouter(store, nil, nil)
|
||||
|
||||
t.Run("empty JSON export", func(t *testing.T) {
|
||||
rr := makeRequest(t, router, "GET", "/api/traces/export?format=json", "1")
|
||||
@@ -400,7 +400,7 @@ func TestMetricsEndpoint(t *testing.T) {
|
||||
metrics.IncError()
|
||||
metrics.SetActiveAgents(3)
|
||||
|
||||
router := NewRouter(store, metrics)
|
||||
router := NewRouter(store, metrics, nil)
|
||||
rr := makeRequest(t, router, "GET", "/metrics", "")
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d", rr.Code)
|
||||
@@ -426,7 +426,7 @@ func TestMetricsEndpoint(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("metrics disabled returns 404", func(t *testing.T) {
|
||||
router := NewRouter(store, nil)
|
||||
router := NewRouter(store, nil, nil)
|
||||
rr := makeRequest(t, router, "GET", "/metrics", "")
|
||||
if rr.Code != http.StatusNotFound {
|
||||
t.Errorf("status = %d, want %d", rr.Code, http.StatusNotFound)
|
||||
@@ -443,7 +443,7 @@ func TestMultiOwnerIsolation(t *testing.T) {
|
||||
seedTraces(t, store, "2", "bob-bot", []string{"send_message", "error", "join_channel"})
|
||||
seedTraces(t, store, "3", "charlie-bot", []string{"send_message"})
|
||||
|
||||
router := NewRouter(store, nil)
|
||||
router := NewRouter(store, nil, nil)
|
||||
|
||||
for _, tc := range []struct {
|
||||
ownerID string
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
package attachments
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// CAS implements content-addressable storage on the local filesystem.
|
||||
// Files are stored at {baseDir}/{hash[0:2]}/{hash[2:4]}/{hash}.
|
||||
type CAS struct {
|
||||
baseDir string
|
||||
logger *slog.Logger
|
||||
mu sync.Mutex // protects concurrent writes of the same hash
|
||||
}
|
||||
|
||||
// NewCAS creates a new content-addressable store rooted at baseDir.
|
||||
// The directory tree is created automatically if it does not exist.
|
||||
func NewCAS(baseDir string, logger *slog.Logger) (*CAS, error) {
|
||||
if err := os.MkdirAll(baseDir, 0o755); err != nil {
|
||||
return nil, fmt.Errorf("create CAS base directory: %w", err)
|
||||
}
|
||||
return &CAS{
|
||||
baseDir: baseDir,
|
||||
logger: logger.With("component", "cas"),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// shardPath returns the full filesystem path for the given hash.
|
||||
func (c *CAS) shardPath(hash string) string {
|
||||
return filepath.Join(c.baseDir, hash[0:2], hash[2:4], hash)
|
||||
}
|
||||
|
||||
// Write streams content from r, computing its SHA-256 hash while writing to
|
||||
// a temporary file. On success the temp file is atomically renamed to the
|
||||
// content-addressed path. If the file already exists (dedup), the temp file
|
||||
// is removed and the existing hash is returned.
|
||||
//
|
||||
// Returns ErrEmptyFile if r yields zero bytes.
|
||||
func (c *CAS) Write(r io.Reader) (hash string, size int64, err error) {
|
||||
// Write to a temp file while computing the hash.
|
||||
tmpFile, err := os.CreateTemp(c.baseDir, ".cas-upload-*")
|
||||
if err != nil {
|
||||
return "", 0, fmt.Errorf("create temp file: %w", err)
|
||||
}
|
||||
tmpPath := tmpFile.Name()
|
||||
|
||||
// Ensure cleanup on any error path.
|
||||
defer func() {
|
||||
if err != nil {
|
||||
tmpFile.Close()
|
||||
os.Remove(tmpPath)
|
||||
}
|
||||
}()
|
||||
|
||||
hasher := sha256.New()
|
||||
w := io.MultiWriter(tmpFile, hasher)
|
||||
|
||||
size, err = io.Copy(w, r)
|
||||
if err != nil {
|
||||
return "", 0, fmt.Errorf("write content: %w", err)
|
||||
}
|
||||
|
||||
if size == 0 {
|
||||
err = ErrEmptyFile
|
||||
return "", 0, err
|
||||
}
|
||||
|
||||
if err = tmpFile.Close(); err != nil {
|
||||
return "", 0, fmt.Errorf("close temp file: %w", err)
|
||||
}
|
||||
|
||||
hash = hex.EncodeToString(hasher.Sum(nil))
|
||||
destPath := c.shardPath(hash)
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
// Check for dedup: if the file already exists, skip the rename.
|
||||
if _, statErr := os.Stat(destPath); statErr == nil {
|
||||
os.Remove(tmpPath)
|
||||
c.logger.Debug("dedup: file already exists", "hash", hash)
|
||||
return hash, size, nil
|
||||
}
|
||||
|
||||
// Create shard directories.
|
||||
destDir := filepath.Dir(destPath)
|
||||
if err = os.MkdirAll(destDir, 0o755); err != nil {
|
||||
return "", 0, fmt.Errorf("create shard directory: %w", err)
|
||||
}
|
||||
|
||||
// Atomic rename.
|
||||
if err = os.Rename(tmpPath, destPath); err != nil {
|
||||
return "", 0, fmt.Errorf("rename temp file: %w", err)
|
||||
}
|
||||
|
||||
c.logger.Info("file stored", "hash", hash, "size", size)
|
||||
return hash, size, nil
|
||||
}
|
||||
|
||||
// Read opens the file identified by hash and returns a ReadCloser.
|
||||
// Returns ErrNotFound if the file does not exist on disk.
|
||||
func (c *CAS) Read(hash string) (io.ReadCloser, error) {
|
||||
path := c.shardPath(hash)
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, fmt.Errorf("open file: %w", err)
|
||||
}
|
||||
return f, nil
|
||||
}
|
||||
|
||||
// Exists returns true if the file identified by hash exists on disk.
|
||||
func (c *CAS) Exists(hash string) bool {
|
||||
_, err := os.Stat(c.shardPath(hash))
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// Delete removes the file identified by hash from disk.
|
||||
// Returns the size of the deleted file, or 0 if the file did not exist.
|
||||
func (c *CAS) Delete(hash string) (int64, error) {
|
||||
path := c.shardPath(hash)
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return 0, nil
|
||||
}
|
||||
return 0, fmt.Errorf("stat file: %w", err)
|
||||
}
|
||||
|
||||
size := info.Size()
|
||||
if err := os.Remove(path); err != nil {
|
||||
return 0, fmt.Errorf("remove file: %w", err)
|
||||
}
|
||||
|
||||
c.logger.Info("file deleted", "hash", hash, "size", size)
|
||||
return size, nil
|
||||
}
|
||||
@@ -0,0 +1,263 @@
|
||||
package attachments
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func newTestCAS(t *testing.T) *CAS {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
cas, err := NewCAS(filepath.Join(dir, "attachments"), slog.Default())
|
||||
if err != nil {
|
||||
t.Fatalf("NewCAS: %v", err)
|
||||
}
|
||||
return cas
|
||||
}
|
||||
|
||||
func sha256Hex(data []byte) string {
|
||||
h := sha256.Sum256(data)
|
||||
return hex.EncodeToString(h[:])
|
||||
}
|
||||
|
||||
func TestCAS_Write(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
content []byte
|
||||
wantErr error
|
||||
}{
|
||||
{
|
||||
name: "normal file",
|
||||
content: []byte("hello world"),
|
||||
},
|
||||
{
|
||||
name: "binary content",
|
||||
content: []byte{0x00, 0x01, 0xff, 0xfe, 0x89, 0x50, 0x4e, 0x47},
|
||||
},
|
||||
{
|
||||
name: "zero bytes",
|
||||
content: []byte{},
|
||||
wantErr: ErrEmptyFile,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cas := newTestCAS(t)
|
||||
hash, size, err := cas.Write(bytes.NewReader(tt.content))
|
||||
|
||||
if tt.wantErr != nil {
|
||||
if err != tt.wantErr {
|
||||
t.Fatalf("expected error %v, got %v", tt.wantErr, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("Write: %v", err)
|
||||
}
|
||||
|
||||
wantHash := sha256Hex(tt.content)
|
||||
if hash != wantHash {
|
||||
t.Errorf("hash = %s, want %s", hash, wantHash)
|
||||
}
|
||||
if size != int64(len(tt.content)) {
|
||||
t.Errorf("size = %d, want %d", size, len(tt.content))
|
||||
}
|
||||
|
||||
// Verify sharded path exists.
|
||||
path := filepath.Join(cas.baseDir, hash[0:2], hash[2:4], hash)
|
||||
if _, err := os.Stat(path); err != nil {
|
||||
t.Errorf("file not found at sharded path: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCAS_Read(t *testing.T) {
|
||||
cas := newTestCAS(t)
|
||||
content := []byte("test content for reading")
|
||||
|
||||
hash, _, err := cas.Write(bytes.NewReader(content))
|
||||
if err != nil {
|
||||
t.Fatalf("Write: %v", err)
|
||||
}
|
||||
|
||||
t.Run("existing file", func(t *testing.T) {
|
||||
rc, err := cas.Read(hash)
|
||||
if err != nil {
|
||||
t.Fatalf("Read: %v", err)
|
||||
}
|
||||
defer rc.Close()
|
||||
|
||||
got, err := io.ReadAll(rc)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadAll: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, content) {
|
||||
t.Errorf("content mismatch: got %q, want %q", got, content)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing file", func(t *testing.T) {
|
||||
_, err := cas.Read("deadbeef" + strings.Repeat("0", 56))
|
||||
if err != ErrNotFound {
|
||||
t.Errorf("expected ErrNotFound, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestCAS_Exists(t *testing.T) {
|
||||
cas := newTestCAS(t)
|
||||
content := []byte("existence check")
|
||||
|
||||
hash, _, err := cas.Write(bytes.NewReader(content))
|
||||
if err != nil {
|
||||
t.Fatalf("Write: %v", err)
|
||||
}
|
||||
|
||||
if !cas.Exists(hash) {
|
||||
t.Error("Exists returned false for stored file")
|
||||
}
|
||||
if cas.Exists("deadbeef" + strings.Repeat("0", 56)) {
|
||||
t.Error("Exists returned true for non-existent file")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCAS_Delete(t *testing.T) {
|
||||
cas := newTestCAS(t)
|
||||
content := []byte("deletable content")
|
||||
|
||||
hash, _, err := cas.Write(bytes.NewReader(content))
|
||||
if err != nil {
|
||||
t.Fatalf("Write: %v", err)
|
||||
}
|
||||
|
||||
t.Run("delete existing", func(t *testing.T) {
|
||||
size, err := cas.Delete(hash)
|
||||
if err != nil {
|
||||
t.Fatalf("Delete: %v", err)
|
||||
}
|
||||
if size != int64(len(content)) {
|
||||
t.Errorf("deleted size = %d, want %d", size, len(content))
|
||||
}
|
||||
if cas.Exists(hash) {
|
||||
t.Error("file still exists after delete")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("delete non-existent", func(t *testing.T) {
|
||||
size, err := cas.Delete("deadbeef" + strings.Repeat("0", 56))
|
||||
if err != nil {
|
||||
t.Fatalf("Delete non-existent: %v", err)
|
||||
}
|
||||
if size != 0 {
|
||||
t.Errorf("deleted size = %d, want 0", size)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestCAS_Dedup(t *testing.T) {
|
||||
cas := newTestCAS(t)
|
||||
content := []byte("duplicate content")
|
||||
|
||||
hash1, _, err := cas.Write(bytes.NewReader(content))
|
||||
if err != nil {
|
||||
t.Fatalf("first Write: %v", err)
|
||||
}
|
||||
|
||||
hash2, _, err := cas.Write(bytes.NewReader(content))
|
||||
if err != nil {
|
||||
t.Fatalf("second Write: %v", err)
|
||||
}
|
||||
|
||||
if hash1 != hash2 {
|
||||
t.Errorf("hashes differ: %s vs %s", hash1, hash2)
|
||||
}
|
||||
|
||||
// Verify only one file on disk.
|
||||
path := filepath.Join(cas.baseDir, hash1[0:2], hash1[2:4], hash1)
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
t.Fatalf("file not found: %v", err)
|
||||
}
|
||||
if info.Size() != int64(len(content)) {
|
||||
t.Errorf("file size = %d, want %d", info.Size(), len(content))
|
||||
}
|
||||
}
|
||||
|
||||
func TestCAS_ConcurrentWrite(t *testing.T) {
|
||||
cas := newTestCAS(t)
|
||||
content := []byte("concurrent content")
|
||||
expected := sha256Hex(content)
|
||||
|
||||
const n = 10
|
||||
var wg sync.WaitGroup
|
||||
errs := make([]error, n)
|
||||
hashes := make([]string, n)
|
||||
|
||||
wg.Add(n)
|
||||
for i := 0; i < n; i++ {
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
h, _, err := cas.Write(bytes.NewReader(content))
|
||||
hashes[idx] = h
|
||||
errs[idx] = err
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
if errs[i] != nil {
|
||||
t.Errorf("goroutine %d: %v", i, errs[i])
|
||||
}
|
||||
if hashes[i] != expected {
|
||||
t.Errorf("goroutine %d: hash = %s, want %s", i, hashes[i], expected)
|
||||
}
|
||||
}
|
||||
|
||||
// Verify single file on disk.
|
||||
if !cas.Exists(expected) {
|
||||
t.Error("file does not exist after concurrent writes")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCAS_DirectoryStructure(t *testing.T) {
|
||||
cas := newTestCAS(t)
|
||||
content := []byte("structure check")
|
||||
|
||||
hash, _, err := cas.Write(bytes.NewReader(content))
|
||||
if err != nil {
|
||||
t.Fatalf("Write: %v", err)
|
||||
}
|
||||
|
||||
// Check two-level sharding.
|
||||
level1 := filepath.Join(cas.baseDir, hash[0:2])
|
||||
level2 := filepath.Join(level1, hash[2:4])
|
||||
file := filepath.Join(level2, hash)
|
||||
|
||||
for _, p := range []string{level1, level2} {
|
||||
info, err := os.Stat(p)
|
||||
if err != nil {
|
||||
t.Fatalf("directory %s not found: %v", p, err)
|
||||
}
|
||||
if !info.IsDir() {
|
||||
t.Errorf("%s is not a directory", p)
|
||||
}
|
||||
}
|
||||
|
||||
info, err := os.Stat(file)
|
||||
if err != nil {
|
||||
t.Fatalf("file not found: %v", err)
|
||||
}
|
||||
if info.IsDir() {
|
||||
t.Error("file is a directory")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
// Package attachments provides content-addressable file storage for SynapBus.
|
||||
//
|
||||
// Files are stored on the local filesystem using a two-level sharded directory
|
||||
// structure based on the SHA-256 hash of the content:
|
||||
//
|
||||
// {data_dir}/attachments/{hash[0:2]}/{hash[2:4]}/{hash}
|
||||
//
|
||||
// Deduplication is automatic: identical content produces the same hash and is
|
||||
// stored only once on disk, while each upload creates its own metadata row in
|
||||
// SQLite. Garbage collection removes orphaned files that are no longer
|
||||
// referenced by any message.
|
||||
package attachments
|
||||
@@ -0,0 +1,120 @@
|
||||
package attachments
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// extensionMIMETypes maps file extensions to MIME types for common types
|
||||
// that net/http.DetectContentType may not identify from magic bytes alone.
|
||||
var extensionMIMETypes = map[string]string{
|
||||
".css": "text/css",
|
||||
".csv": "text/csv",
|
||||
".doc": "application/msword",
|
||||
".docx": "application/vnd.openxmlformats-officedocument.wordprocessingml.document",
|
||||
".gif": "image/gif",
|
||||
".html": "text/html",
|
||||
".jpeg": "image/jpeg",
|
||||
".jpg": "image/jpeg",
|
||||
".js": "application/javascript",
|
||||
".json": "application/json",
|
||||
".md": "text/markdown",
|
||||
".mp3": "audio/mpeg",
|
||||
".mp4": "video/mp4",
|
||||
".pdf": "application/pdf",
|
||||
".png": "image/png",
|
||||
".pptx": "application/vnd.openxmlformats-officedocument.presentationml.presentation",
|
||||
".svg": "image/svg+xml",
|
||||
".tar": "application/x-tar",
|
||||
".txt": "text/plain",
|
||||
".wav": "audio/wav",
|
||||
".webp": "image/webp",
|
||||
".xlsx": "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
|
||||
".xml": "application/xml",
|
||||
".yaml": "text/yaml",
|
||||
".yml": "text/yaml",
|
||||
".zip": "application/zip",
|
||||
}
|
||||
|
||||
// imageTypes is the set of MIME types considered "image" for inline preview.
|
||||
var imageTypes = map[string]bool{
|
||||
"image/jpeg": true,
|
||||
"image/png": true,
|
||||
"image/gif": true,
|
||||
"image/webp": true,
|
||||
"image/svg+xml": true,
|
||||
}
|
||||
|
||||
// DetectMIMEType detects the MIME type from file content (magic bytes) with
|
||||
// fallback to extension-based detection. Returns "application/octet-stream"
|
||||
// if detection fails entirely.
|
||||
//
|
||||
// When magic-byte detection returns a generic type like "text/plain" or
|
||||
// "application/octet-stream", the extension is checked for a more specific
|
||||
// type (e.g., .csv -> text/csv, .json -> application/json).
|
||||
func DetectMIMEType(content []byte, filename string) string {
|
||||
var detected string
|
||||
if len(content) > 0 {
|
||||
detected = http.DetectContentType(content)
|
||||
}
|
||||
|
||||
// If detection returned a specific, non-generic type, use it.
|
||||
if detected != "" && !isGenericMIME(detected) {
|
||||
return detected
|
||||
}
|
||||
|
||||
// Extension-based lookup for more specificity.
|
||||
if filename != "" {
|
||||
ext := strings.ToLower(filepath.Ext(filename))
|
||||
if mime, ok := extensionMIMETypes[ext]; ok {
|
||||
return mime
|
||||
}
|
||||
}
|
||||
|
||||
// Return whatever detection found, or fall back to octet-stream.
|
||||
if detected != "" {
|
||||
return detected
|
||||
}
|
||||
return "application/octet-stream"
|
||||
}
|
||||
|
||||
// isGenericMIME returns true if the MIME type is a generic catch-all that
|
||||
// should be overridden by extension-based detection when available.
|
||||
func isGenericMIME(mimeType string) bool {
|
||||
// Normalize: strip parameters like "; charset=utf-8".
|
||||
base := mimeType
|
||||
if idx := strings.Index(mimeType, ";"); idx >= 0 {
|
||||
base = strings.TrimSpace(mimeType[:idx])
|
||||
}
|
||||
return base == "application/octet-stream" || base == "text/plain"
|
||||
}
|
||||
|
||||
// DefaultFilename returns a default filename based on MIME type when the
|
||||
// caller provides none.
|
||||
func DefaultFilename(mimeType string) string {
|
||||
switch {
|
||||
case strings.HasPrefix(mimeType, "image/png"):
|
||||
return "untitled.png"
|
||||
case strings.HasPrefix(mimeType, "image/jpeg"):
|
||||
return "untitled.jpg"
|
||||
case strings.HasPrefix(mimeType, "image/gif"):
|
||||
return "untitled.gif"
|
||||
case strings.HasPrefix(mimeType, "image/webp"):
|
||||
return "untitled.webp"
|
||||
case strings.HasPrefix(mimeType, "image/svg"):
|
||||
return "untitled.svg"
|
||||
case mimeType == "application/pdf":
|
||||
return "untitled.pdf"
|
||||
case strings.HasPrefix(mimeType, "text/"):
|
||||
return "untitled.txt"
|
||||
default:
|
||||
return "untitled.bin"
|
||||
}
|
||||
}
|
||||
|
||||
// IsImageType returns true if the MIME type is a supported image type for
|
||||
// inline preview in the Web UI.
|
||||
func IsImageType(mimeType string) bool {
|
||||
return imageTypes[mimeType]
|
||||
}
|
||||
@@ -0,0 +1,162 @@
|
||||
package attachments
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestDetectMIMEType(t *testing.T) {
|
||||
// PNG magic bytes.
|
||||
pngHeader := []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a}
|
||||
// JPEG magic bytes.
|
||||
jpegHeader := []byte{0xff, 0xd8, 0xff, 0xe0}
|
||||
// GIF magic bytes.
|
||||
gifHeader := []byte("GIF89a")
|
||||
// PDF magic bytes.
|
||||
pdfHeader := []byte("%PDF-1.4")
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
content []byte
|
||||
filename string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "PNG from magic bytes",
|
||||
content: pngHeader,
|
||||
filename: "",
|
||||
want: "image/png",
|
||||
},
|
||||
{
|
||||
name: "JPEG from magic bytes",
|
||||
content: jpegHeader,
|
||||
filename: "",
|
||||
want: "image/jpeg",
|
||||
},
|
||||
{
|
||||
name: "GIF from magic bytes",
|
||||
content: gifHeader,
|
||||
filename: "",
|
||||
want: "image/gif",
|
||||
},
|
||||
{
|
||||
name: "PDF from magic bytes",
|
||||
content: pdfHeader,
|
||||
filename: "report.pdf",
|
||||
want: "application/pdf",
|
||||
},
|
||||
{
|
||||
name: "CSV from extension",
|
||||
content: []byte("a,b,c\n1,2,3"),
|
||||
filename: "data.csv",
|
||||
want: "text/csv",
|
||||
},
|
||||
{
|
||||
name: "JSON from extension",
|
||||
content: []byte(`{"key": "value"}`),
|
||||
filename: "config.json",
|
||||
want: "application/json",
|
||||
},
|
||||
{
|
||||
name: "Markdown from extension",
|
||||
content: []byte("# Hello"),
|
||||
filename: "readme.md",
|
||||
want: "text/markdown",
|
||||
},
|
||||
{
|
||||
name: "ZIP from extension",
|
||||
content: []byte("not real zip"),
|
||||
filename: "archive.zip",
|
||||
want: "application/zip",
|
||||
},
|
||||
{
|
||||
name: "unknown falls back to octet-stream",
|
||||
content: []byte{0x00, 0x01, 0x02},
|
||||
filename: "mystery.xyz",
|
||||
want: "application/octet-stream",
|
||||
},
|
||||
{
|
||||
name: "empty content with extension",
|
||||
content: nil,
|
||||
filename: "test.png",
|
||||
want: "image/png",
|
||||
},
|
||||
{
|
||||
name: "empty everything",
|
||||
content: nil,
|
||||
filename: "",
|
||||
want: "application/octet-stream",
|
||||
},
|
||||
{
|
||||
name: "HTML from magic bytes",
|
||||
content: []byte("<html><body>hello</body></html>"),
|
||||
filename: "",
|
||||
want: "text/html; charset=utf-8",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := DetectMIMEType(tt.content, tt.filename)
|
||||
if got != tt.want {
|
||||
t.Errorf("DetectMIMEType(%v, %q) = %q, want %q", tt.content[:min(len(tt.content), 8)], tt.filename, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultFilename(t *testing.T) {
|
||||
tests := []struct {
|
||||
mimeType string
|
||||
want string
|
||||
}{
|
||||
{"image/png", "untitled.png"},
|
||||
{"image/jpeg", "untitled.jpg"},
|
||||
{"image/gif", "untitled.gif"},
|
||||
{"image/webp", "untitled.webp"},
|
||||
{"image/svg+xml", "untitled.svg"},
|
||||
{"application/pdf", "untitled.pdf"},
|
||||
{"text/plain", "untitled.txt"},
|
||||
{"text/csv", "untitled.txt"},
|
||||
{"application/octet-stream", "untitled.bin"},
|
||||
{"application/zip", "untitled.bin"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.mimeType, func(t *testing.T) {
|
||||
got := DefaultFilename(tt.mimeType)
|
||||
if got != tt.want {
|
||||
t.Errorf("DefaultFilename(%q) = %q, want %q", tt.mimeType, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsImageType(t *testing.T) {
|
||||
tests := []struct {
|
||||
mimeType string
|
||||
want bool
|
||||
}{
|
||||
{"image/jpeg", true},
|
||||
{"image/png", true},
|
||||
{"image/gif", true},
|
||||
{"image/webp", true},
|
||||
{"image/svg+xml", true},
|
||||
{"application/pdf", false},
|
||||
{"text/plain", false},
|
||||
{"image/tiff", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.mimeType, func(t *testing.T) {
|
||||
got := IsImageType(tt.mimeType)
|
||||
if got != tt.want {
|
||||
t.Errorf("IsImageType(%q) = %v, want %v", tt.mimeType, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func min(a, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
package attachments
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"time"
|
||||
)
|
||||
|
||||
// MaxFileSize is the maximum allowed upload size (50 MB).
|
||||
const MaxFileSize = 50 * 1024 * 1024 // 50 MB
|
||||
|
||||
// Sentinel errors.
|
||||
var (
|
||||
ErrNotFound = errors.New("attachment not found")
|
||||
ErrFileTooLarge = errors.New("file exceeds maximum size of 50MB")
|
||||
ErrEmptyFile = errors.New("empty file not allowed")
|
||||
ErrFileMissing = errors.New("attachment file missing from disk")
|
||||
)
|
||||
|
||||
// Attachment represents the metadata for a stored file.
|
||||
type Attachment struct {
|
||||
ID int64 `json:"id"`
|
||||
Hash string `json:"hash"`
|
||||
OriginalFilename string `json:"original_filename"`
|
||||
Size int64 `json:"size"`
|
||||
MIMEType string `json:"mime_type"`
|
||||
MessageID *int64 `json:"message_id,omitempty"`
|
||||
UploadedBy string `json:"uploaded_by"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
// UploadRequest contains the parameters for uploading an attachment.
|
||||
type UploadRequest struct {
|
||||
Content io.Reader
|
||||
Filename string
|
||||
MIMEType string
|
||||
MessageID *int64
|
||||
UploadedBy string
|
||||
}
|
||||
|
||||
// UploadResult contains the result of a successful upload.
|
||||
type UploadResult struct {
|
||||
Hash string `json:"hash"`
|
||||
Size int64 `json:"size"`
|
||||
MIMEType string `json:"mime_type"`
|
||||
Filename string `json:"original_filename"`
|
||||
}
|
||||
|
||||
// DownloadResult contains the result of a successful download.
|
||||
type DownloadResult struct {
|
||||
Content io.ReadCloser
|
||||
Hash string
|
||||
Filename string
|
||||
MIMEType string
|
||||
Size int64
|
||||
}
|
||||
|
||||
// GCResult contains the result of a garbage collection run.
|
||||
type GCResult struct {
|
||||
FilesRemoved int `json:"files_removed"`
|
||||
BytesReclaimed int64 `json:"bytes_reclaimed"`
|
||||
}
|
||||
@@ -0,0 +1,224 @@
|
||||
package attachments
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Service is the main entry point for all attachment operations.
|
||||
// It composes the metadata Store with the CAS filesystem engine.
|
||||
type Service struct {
|
||||
store Store
|
||||
cas *CAS
|
||||
logger *slog.Logger
|
||||
gcMu sync.Mutex // prevents concurrent GC + upload conflicts
|
||||
}
|
||||
|
||||
// NewService creates a new attachment service.
|
||||
func NewService(store Store, cas *CAS, logger *slog.Logger) *Service {
|
||||
return &Service{
|
||||
store: store,
|
||||
cas: cas,
|
||||
logger: logger.With("component", "attachment-service"),
|
||||
}
|
||||
}
|
||||
|
||||
// Upload stores a file in content-addressable storage and records metadata.
|
||||
// It validates size, detects MIME type, deduplicates on hash, and logs the
|
||||
// operation.
|
||||
func (s *Service) Upload(ctx context.Context, req UploadRequest) (*UploadResult, error) {
|
||||
start := time.Now()
|
||||
|
||||
// Read the content into a limited reader to enforce size limits.
|
||||
// We read up to MaxFileSize + 1 to detect overflow.
|
||||
content, err := io.ReadAll(io.LimitReader(req.Content, MaxFileSize+1))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read content: %w", err)
|
||||
}
|
||||
|
||||
if len(content) == 0 {
|
||||
return nil, ErrEmptyFile
|
||||
}
|
||||
|
||||
if int64(len(content)) > MaxFileSize {
|
||||
return nil, ErrFileTooLarge
|
||||
}
|
||||
|
||||
// Detect MIME type if not provided.
|
||||
mimeType := req.MIMEType
|
||||
if mimeType == "" {
|
||||
// Use first 512 bytes for detection.
|
||||
sniffBuf := content
|
||||
if len(sniffBuf) > 512 {
|
||||
sniffBuf = sniffBuf[:512]
|
||||
}
|
||||
mimeType = DetectMIMEType(sniffBuf, req.Filename)
|
||||
}
|
||||
|
||||
// Assign default filename if missing.
|
||||
filename := req.Filename
|
||||
if filename == "" {
|
||||
filename = DefaultFilename(mimeType)
|
||||
}
|
||||
|
||||
// Write to CAS.
|
||||
hash, size, err := s.cas.Write(bytes.NewReader(content))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("write to CAS: %w", err)
|
||||
}
|
||||
|
||||
// Store metadata.
|
||||
att := &Attachment{
|
||||
Hash: hash,
|
||||
OriginalFilename: filename,
|
||||
Size: size,
|
||||
MIMEType: mimeType,
|
||||
MessageID: req.MessageID,
|
||||
UploadedBy: req.UploadedBy,
|
||||
}
|
||||
|
||||
if err := s.store.InsertMetadata(ctx, att); err != nil {
|
||||
return nil, fmt.Errorf("insert metadata: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("attachment uploaded",
|
||||
"hash", hash,
|
||||
"filename", filename,
|
||||
"size", size,
|
||||
"mime_type", mimeType,
|
||||
"uploaded_by", req.UploadedBy,
|
||||
"duration_ms", time.Since(start).Milliseconds(),
|
||||
)
|
||||
|
||||
return &UploadResult{
|
||||
Hash: hash,
|
||||
Size: size,
|
||||
MIMEType: mimeType,
|
||||
Filename: filename,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Download retrieves a file and its metadata by hash.
|
||||
// Returns ErrNotFound if no metadata exists, or ErrFileMissing if the
|
||||
// metadata exists but the file is not on disk.
|
||||
func (s *Service) Download(ctx context.Context, hash string) (*DownloadResult, error) {
|
||||
start := time.Now()
|
||||
|
||||
// Look up metadata.
|
||||
atts, err := s.store.GetByHash(ctx, hash)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get metadata: %w", err)
|
||||
}
|
||||
if len(atts) == 0 {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
|
||||
// Use the first metadata row for filename/mime info.
|
||||
meta := atts[0]
|
||||
|
||||
// Open the file from CAS.
|
||||
reader, err := s.cas.Read(hash)
|
||||
if err != nil {
|
||||
if err == ErrNotFound {
|
||||
s.logger.Warn("attachment file missing from disk",
|
||||
"hash", hash,
|
||||
"filename", meta.OriginalFilename,
|
||||
)
|
||||
return nil, ErrFileMissing
|
||||
}
|
||||
return nil, fmt.Errorf("read from CAS: %w", err)
|
||||
}
|
||||
|
||||
s.logger.Info("attachment downloaded",
|
||||
"hash", hash,
|
||||
"filename", meta.OriginalFilename,
|
||||
"size", meta.Size,
|
||||
"duration_ms", time.Since(start).Milliseconds(),
|
||||
)
|
||||
|
||||
return &DownloadResult{
|
||||
Content: reader,
|
||||
Hash: hash,
|
||||
Filename: meta.OriginalFilename,
|
||||
MIMEType: meta.MIMEType,
|
||||
Size: meta.Size,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// AttachToMessage links an existing attachment to a message by updating
|
||||
// the message_id on all metadata rows with the given hash that are currently
|
||||
// unlinked.
|
||||
func (s *Service) AttachToMessage(ctx context.Context, hash string, messageID int64) error {
|
||||
atts, err := s.store.GetByHash(ctx, hash)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get metadata: %w", err)
|
||||
}
|
||||
if len(atts) == 0 {
|
||||
return ErrNotFound
|
||||
}
|
||||
|
||||
// We update by inserting a new metadata row linked to the message,
|
||||
// copying from the first existing row.
|
||||
meta := atts[0]
|
||||
linked := &Attachment{
|
||||
Hash: meta.Hash,
|
||||
OriginalFilename: meta.OriginalFilename,
|
||||
Size: meta.Size,
|
||||
MIMEType: meta.MIMEType,
|
||||
MessageID: &messageID,
|
||||
UploadedBy: meta.UploadedBy,
|
||||
}
|
||||
return s.store.InsertMetadata(ctx, linked)
|
||||
}
|
||||
|
||||
// GetByMessageID returns all attachments for a given message.
|
||||
func (s *Service) GetByMessageID(ctx context.Context, messageID int64) ([]*Attachment, error) {
|
||||
return s.store.GetByMessageID(ctx, messageID)
|
||||
}
|
||||
|
||||
// GarbageCollect finds and removes orphaned attachments — files that are
|
||||
// no longer referenced by any message. Returns a summary of what was removed.
|
||||
func (s *Service) GarbageCollect(ctx context.Context) (*GCResult, error) {
|
||||
s.gcMu.Lock()
|
||||
defer s.gcMu.Unlock()
|
||||
|
||||
start := time.Now()
|
||||
|
||||
orphans, err := s.store.FindOrphanHashes(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("find orphans: %w", err)
|
||||
}
|
||||
|
||||
result := &GCResult{}
|
||||
|
||||
for _, hash := range orphans {
|
||||
// Delete from CAS.
|
||||
size, err := s.cas.Delete(hash)
|
||||
if err != nil {
|
||||
s.logger.Error("GC: failed to delete file", "hash", hash, "error", err)
|
||||
continue
|
||||
}
|
||||
|
||||
// Delete metadata.
|
||||
if err := s.store.DeleteByHash(ctx, hash); err != nil {
|
||||
s.logger.Error("GC: failed to delete metadata", "hash", hash, "error", err)
|
||||
continue
|
||||
}
|
||||
|
||||
result.FilesRemoved++
|
||||
result.BytesReclaimed += size
|
||||
}
|
||||
|
||||
s.logger.Info("garbage collection complete",
|
||||
"files_removed", result.FilesRemoved,
|
||||
"bytes_reclaimed", result.BytesReclaimed,
|
||||
"duration_ms", time.Since(start).Milliseconds(),
|
||||
)
|
||||
|
||||
return result, nil
|
||||
}
|
||||
@@ -0,0 +1,265 @@
|
||||
package attachments
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"io"
|
||||
"log/slog"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func newTestService(t *testing.T) (*Service, *sql.DB) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
dir := t.TempDir()
|
||||
|
||||
cas, err := NewCAS(filepath.Join(dir, "attachments"), slog.Default())
|
||||
if err != nil {
|
||||
t.Fatalf("NewCAS: %v", err)
|
||||
}
|
||||
|
||||
store := NewSQLiteStore(db, slog.Default())
|
||||
svc := NewService(store, cas, slog.Default())
|
||||
return svc, db
|
||||
}
|
||||
|
||||
func TestService_Upload(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("service upload test content")
|
||||
wantHash := sha256Hex(content)
|
||||
|
||||
result, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader(content),
|
||||
Filename: "test.txt",
|
||||
UploadedBy: "agent-a",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Upload: %v", err)
|
||||
}
|
||||
|
||||
if result.Hash != wantHash {
|
||||
t.Errorf("hash = %s, want %s", result.Hash, wantHash)
|
||||
}
|
||||
if result.Size != int64(len(content)) {
|
||||
t.Errorf("size = %d, want %d", result.Size, len(content))
|
||||
}
|
||||
if result.Filename != "test.txt" {
|
||||
t.Errorf("filename = %s, want test.txt", result.Filename)
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_Upload_MIMEDetection(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// PNG magic bytes.
|
||||
pngContent := []byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, 0x00, 0x00}
|
||||
|
||||
result, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader(pngContent),
|
||||
Filename: "image.png",
|
||||
UploadedBy: "agent-a",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Upload: %v", err)
|
||||
}
|
||||
|
||||
if result.MIMEType != "image/png" {
|
||||
t.Errorf("mime_type = %s, want image/png", result.MIMEType)
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_Upload_DefaultFilename(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("no filename content here")
|
||||
|
||||
result, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader(content),
|
||||
UploadedBy: "agent-a",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Upload: %v", err)
|
||||
}
|
||||
|
||||
if result.Filename == "" {
|
||||
t.Error("expected a default filename, got empty")
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_Upload_EmptyFile(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader([]byte{}),
|
||||
Filename: "empty.txt",
|
||||
UploadedBy: "agent-a",
|
||||
})
|
||||
if err != ErrEmptyFile {
|
||||
t.Errorf("expected ErrEmptyFile, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_Upload_SizeLimit(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Create a reader that yields MaxFileSize+1 bytes.
|
||||
oversized := make([]byte, MaxFileSize+1)
|
||||
for i := range oversized {
|
||||
oversized[i] = 'x'
|
||||
}
|
||||
|
||||
_, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader(oversized),
|
||||
Filename: "toobig.bin",
|
||||
UploadedBy: "agent-a",
|
||||
})
|
||||
if err != ErrFileTooLarge {
|
||||
t.Errorf("expected ErrFileTooLarge, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_Download(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("download test content")
|
||||
|
||||
uploadResult, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader(content),
|
||||
Filename: "download.txt",
|
||||
UploadedBy: "agent-a",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Upload: %v", err)
|
||||
}
|
||||
|
||||
t.Run("successful download", func(t *testing.T) {
|
||||
result, err := svc.Download(ctx, uploadResult.Hash)
|
||||
if err != nil {
|
||||
t.Fatalf("Download: %v", err)
|
||||
}
|
||||
defer result.Content.Close()
|
||||
|
||||
got, err := io.ReadAll(result.Content)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadAll: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, content) {
|
||||
t.Error("downloaded content does not match uploaded content")
|
||||
}
|
||||
if result.Filename != "download.txt" {
|
||||
t.Errorf("filename = %s, want download.txt", result.Filename)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("not found", func(t *testing.T) {
|
||||
_, err := svc.Download(ctx, "deadbeef"+strings.Repeat("0", 56))
|
||||
if err != ErrNotFound {
|
||||
t.Errorf("expected ErrNotFound, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestService_Dedup(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
content := []byte("dedup test content")
|
||||
wantHash := sha256Hex(content)
|
||||
|
||||
r1, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader(content),
|
||||
Filename: "version1.txt",
|
||||
UploadedBy: "agent-a",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Upload 1: %v", err)
|
||||
}
|
||||
|
||||
r2, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader(content),
|
||||
Filename: "version2.txt",
|
||||
UploadedBy: "agent-b",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Upload 2: %v", err)
|
||||
}
|
||||
|
||||
if r1.Hash != r2.Hash {
|
||||
t.Errorf("hashes differ: %s vs %s", r1.Hash, r2.Hash)
|
||||
}
|
||||
if r1.Hash != wantHash {
|
||||
t.Errorf("hash = %s, want %s", r1.Hash, wantHash)
|
||||
}
|
||||
}
|
||||
|
||||
func TestService_GarbageCollect(t *testing.T) {
|
||||
svc, db := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Upload an orphan (no message_id).
|
||||
orphanContent := []byte("orphan gc content")
|
||||
_, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader(orphanContent),
|
||||
Filename: "orphan.txt",
|
||||
UploadedBy: "agent-a",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Upload orphan: %v", err)
|
||||
}
|
||||
|
||||
// Upload a linked attachment.
|
||||
msgID := seedMessage(t, db)
|
||||
linkedContent := []byte("linked gc content")
|
||||
linkedResult, err := svc.Upload(ctx, UploadRequest{
|
||||
Content: bytes.NewReader(linkedContent),
|
||||
Filename: "linked.txt",
|
||||
MessageID: &msgID,
|
||||
UploadedBy: "agent-a",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Upload linked: %v", err)
|
||||
}
|
||||
|
||||
// Run GC.
|
||||
gc, err := svc.GarbageCollect(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("GarbageCollect: %v", err)
|
||||
}
|
||||
|
||||
if gc.FilesRemoved != 1 {
|
||||
t.Errorf("files_removed = %d, want 1", gc.FilesRemoved)
|
||||
}
|
||||
|
||||
// Verify linked file still downloadable.
|
||||
result, err := svc.Download(ctx, linkedResult.Hash)
|
||||
if err != nil {
|
||||
t.Fatalf("Download linked after GC: %v", err)
|
||||
}
|
||||
result.Content.Close()
|
||||
}
|
||||
|
||||
func TestService_GarbageCollect_Empty(t *testing.T) {
|
||||
svc, _ := newTestService(t)
|
||||
ctx := context.Background()
|
||||
|
||||
gc, err := svc.GarbageCollect(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("GarbageCollect: %v", err)
|
||||
}
|
||||
if gc.FilesRemoved != 0 {
|
||||
t.Errorf("files_removed = %d, want 0", gc.FilesRemoved)
|
||||
}
|
||||
if gc.BytesReclaimed != 0 {
|
||||
t.Errorf("bytes_reclaimed = %d, want 0", gc.BytesReclaimed)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
package attachments
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
)
|
||||
|
||||
// SQLiteStore implements Store backed by modernc.org/sqlite.
|
||||
type SQLiteStore struct {
|
||||
db *sql.DB
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewSQLiteStore creates a new SQLite-backed attachment metadata store.
|
||||
func NewSQLiteStore(db *sql.DB, logger *slog.Logger) *SQLiteStore {
|
||||
return &SQLiteStore{
|
||||
db: db,
|
||||
logger: logger.With("component", "attachment-store"),
|
||||
}
|
||||
}
|
||||
|
||||
// InsertMetadata inserts a new attachment metadata row.
|
||||
func (s *SQLiteStore) InsertMetadata(ctx context.Context, a *Attachment) error {
|
||||
const query = `INSERT INTO attachments (hash, original_filename, size, mime_type, message_id, uploaded_by, created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)`
|
||||
|
||||
now := time.Now().UTC()
|
||||
result, err := s.db.ExecContext(ctx, query,
|
||||
a.Hash,
|
||||
a.OriginalFilename,
|
||||
a.Size,
|
||||
a.MIMEType,
|
||||
a.MessageID,
|
||||
a.UploadedBy,
|
||||
now,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("insert attachment metadata: %w", err)
|
||||
}
|
||||
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get last insert id: %w", err)
|
||||
}
|
||||
|
||||
a.ID = id
|
||||
a.CreatedAt = now
|
||||
|
||||
s.logger.Debug("attachment metadata inserted",
|
||||
"id", id,
|
||||
"hash", a.Hash,
|
||||
"filename", a.OriginalFilename,
|
||||
)
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetByHash returns all attachment metadata rows matching the given hash.
|
||||
func (s *SQLiteStore) GetByHash(ctx context.Context, hash string) ([]*Attachment, error) {
|
||||
const query = `SELECT id, hash, original_filename, size, mime_type, message_id, uploaded_by, created_at
|
||||
FROM attachments WHERE hash = ? ORDER BY created_at DESC`
|
||||
|
||||
rows, err := s.db.QueryContext(ctx, query, hash)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("query attachments by hash: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return scanAttachments(rows)
|
||||
}
|
||||
|
||||
// GetByMessageID returns all attachment metadata rows for a given message.
|
||||
func (s *SQLiteStore) GetByMessageID(ctx context.Context, messageID int64) ([]*Attachment, error) {
|
||||
const query = `SELECT id, hash, original_filename, size, mime_type, message_id, uploaded_by, created_at
|
||||
FROM attachments WHERE message_id = ? ORDER BY created_at ASC`
|
||||
|
||||
rows, err := s.db.QueryContext(ctx, query, messageID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("query attachments by message_id: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
return scanAttachments(rows)
|
||||
}
|
||||
|
||||
// DeleteByHash removes all metadata rows matching the given hash.
|
||||
func (s *SQLiteStore) DeleteByHash(ctx context.Context, hash string) error {
|
||||
const query = `DELETE FROM attachments WHERE hash = ?`
|
||||
|
||||
result, err := s.db.ExecContext(ctx, query, hash)
|
||||
if err != nil {
|
||||
return fmt.Errorf("delete attachment metadata: %w", err)
|
||||
}
|
||||
|
||||
n, _ := result.RowsAffected()
|
||||
s.logger.Debug("attachment metadata deleted", "hash", hash, "rows", n)
|
||||
return nil
|
||||
}
|
||||
|
||||
// FindOrphanHashes returns hashes that have no valid message reference.
|
||||
// A hash is orphaned if all its metadata rows have message_id IS NULL or
|
||||
// the referenced message no longer exists.
|
||||
func (s *SQLiteStore) FindOrphanHashes(ctx context.Context) ([]string, error) {
|
||||
const query = `SELECT DISTINCT a.hash FROM attachments a
|
||||
WHERE a.message_id IS NULL
|
||||
OR a.message_id NOT IN (SELECT id FROM messages)
|
||||
GROUP BY a.hash
|
||||
HAVING COUNT(CASE WHEN a.message_id IN (SELECT id FROM messages) THEN 1 END) = 0`
|
||||
|
||||
rows, err := s.db.QueryContext(ctx, query)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("find orphan hashes: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var hashes []string
|
||||
for rows.Next() {
|
||||
var h string
|
||||
if err := rows.Scan(&h); err != nil {
|
||||
return nil, fmt.Errorf("scan orphan hash: %w", err)
|
||||
}
|
||||
hashes = append(hashes, h)
|
||||
}
|
||||
return hashes, rows.Err()
|
||||
}
|
||||
|
||||
// CountReferences returns the number of metadata rows referencing the hash.
|
||||
func (s *SQLiteStore) CountReferences(ctx context.Context, hash string) (int64, error) {
|
||||
const query = `SELECT COUNT(*) FROM attachments WHERE hash = ?`
|
||||
|
||||
var count int64
|
||||
if err := s.db.QueryRowContext(ctx, query, hash).Scan(&count); err != nil {
|
||||
return 0, fmt.Errorf("count references: %w", err)
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// scanAttachments scans rows into a slice of Attachment pointers.
|
||||
func scanAttachments(rows *sql.Rows) ([]*Attachment, error) {
|
||||
var attachments []*Attachment
|
||||
for rows.Next() {
|
||||
a := &Attachment{}
|
||||
var createdAtStr string
|
||||
if err := rows.Scan(
|
||||
&a.ID,
|
||||
&a.Hash,
|
||||
&a.OriginalFilename,
|
||||
&a.Size,
|
||||
&a.MIMEType,
|
||||
&a.MessageID,
|
||||
&a.UploadedBy,
|
||||
&createdAtStr,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("scan attachment: %w", err)
|
||||
}
|
||||
// Parse the created_at timestamp.
|
||||
if t, err := time.Parse("2006-01-02 15:04:05", createdAtStr); err == nil {
|
||||
a.CreatedAt = t
|
||||
} else if t, err := time.Parse(time.RFC3339, createdAtStr); err == nil {
|
||||
a.CreatedAt = t
|
||||
}
|
||||
attachments = append(attachments, a)
|
||||
}
|
||||
return attachments, rows.Err()
|
||||
}
|
||||
@@ -0,0 +1,226 @@
|
||||
package attachments
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"testing"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/storage"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
t.Fatalf("open database: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { db.Close() })
|
||||
|
||||
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
|
||||
t.Fatalf("enable foreign keys: %v", err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
if err := storage.RunMigrations(ctx, db); err != nil {
|
||||
t.Fatalf("run migrations: %v", err)
|
||||
}
|
||||
|
||||
// Seed test user and agent.
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`)
|
||||
|
||||
return db
|
||||
}
|
||||
|
||||
func newTestStore(t *testing.T) (*SQLiteStore, *sql.DB) {
|
||||
t.Helper()
|
||||
db := newTestDB(t)
|
||||
store := NewSQLiteStore(db, slog.Default())
|
||||
return store, db
|
||||
}
|
||||
|
||||
// seedMessage inserts a test message so we can reference it via foreign key.
|
||||
func seedMessage(t *testing.T, db *sql.DB) int64 {
|
||||
t.Helper()
|
||||
// Ensure user exists.
|
||||
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`)
|
||||
// Ensure agents exist.
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (id, name, display_name, owner_id, api_key_hash) VALUES (1, 'agent-a', 'Agent A', 1, 'hash')`)
|
||||
db.Exec(`INSERT OR IGNORE INTO agents (id, name, display_name, owner_id, api_key_hash) VALUES (2, 'agent-b', 'Agent B', 1, 'hash')`)
|
||||
// Ensure conversation exists.
|
||||
db.Exec(`INSERT OR IGNORE INTO conversations (id, subject, created_by) VALUES (1, 'test', 'agent-a')`)
|
||||
// Insert message.
|
||||
result, err := db.Exec(`INSERT INTO messages (conversation_id, from_agent, to_agent, body) VALUES (1, 'agent-a', 'agent-b', 'hello')`)
|
||||
if err != nil {
|
||||
t.Fatalf("seed message: %v", err)
|
||||
}
|
||||
id, _ := result.LastInsertId()
|
||||
return id
|
||||
}
|
||||
|
||||
func TestSQLiteStore_InsertMetadata(t *testing.T) {
|
||||
store, _ := newTestStore(t)
|
||||
ctx := context.Background()
|
||||
|
||||
a := &Attachment{
|
||||
Hash: "abcdef1234567890abcdef1234567890abcdef1234567890abcdef1234567890",
|
||||
OriginalFilename: "test.png",
|
||||
Size: 1024,
|
||||
MIMEType: "image/png",
|
||||
UploadedBy: "agent-a",
|
||||
}
|
||||
|
||||
if err := store.InsertMetadata(ctx, a); err != nil {
|
||||
t.Fatalf("InsertMetadata: %v", err)
|
||||
}
|
||||
|
||||
if a.ID == 0 {
|
||||
t.Error("ID not populated after insert")
|
||||
}
|
||||
if a.CreatedAt.IsZero() {
|
||||
t.Error("CreatedAt not populated after insert")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_GetByHash(t *testing.T) {
|
||||
store, _ := newTestStore(t)
|
||||
ctx := context.Background()
|
||||
|
||||
hash := "abcdef1234567890abcdef1234567890abcdef1234567890abcdef1234567890"
|
||||
|
||||
// Insert two rows with the same hash.
|
||||
a1 := &Attachment{Hash: hash, OriginalFilename: "file1.png", Size: 100, MIMEType: "image/png", UploadedBy: "agent-a"}
|
||||
a2 := &Attachment{Hash: hash, OriginalFilename: "file2.png", Size: 100, MIMEType: "image/png", UploadedBy: "agent-b"}
|
||||
|
||||
if err := store.InsertMetadata(ctx, a1); err != nil {
|
||||
t.Fatalf("InsertMetadata a1: %v", err)
|
||||
}
|
||||
if err := store.InsertMetadata(ctx, a2); err != nil {
|
||||
t.Fatalf("InsertMetadata a2: %v", err)
|
||||
}
|
||||
|
||||
results, err := store.GetByHash(ctx, hash)
|
||||
if err != nil {
|
||||
t.Fatalf("GetByHash: %v", err)
|
||||
}
|
||||
if len(results) != 2 {
|
||||
t.Fatalf("expected 2 results, got %d", len(results))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_GetByMessageID(t *testing.T) {
|
||||
store, db := newTestStore(t)
|
||||
ctx := context.Background()
|
||||
|
||||
msgID := seedMessage(t, db)
|
||||
|
||||
a := &Attachment{
|
||||
Hash: "1111111111111111111111111111111111111111111111111111111111111111",
|
||||
OriginalFilename: "report.pdf",
|
||||
Size: 2048,
|
||||
MIMEType: "application/pdf",
|
||||
MessageID: &msgID,
|
||||
UploadedBy: "agent-a",
|
||||
}
|
||||
if err := store.InsertMetadata(ctx, a); err != nil {
|
||||
t.Fatalf("InsertMetadata: %v", err)
|
||||
}
|
||||
|
||||
results, err := store.GetByMessageID(ctx, msgID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetByMessageID: %v", err)
|
||||
}
|
||||
if len(results) != 1 {
|
||||
t.Fatalf("expected 1 result, got %d", len(results))
|
||||
}
|
||||
if results[0].OriginalFilename != "report.pdf" {
|
||||
t.Errorf("filename = %s, want report.pdf", results[0].OriginalFilename)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_DeleteByHash(t *testing.T) {
|
||||
store, _ := newTestStore(t)
|
||||
ctx := context.Background()
|
||||
|
||||
hash := "2222222222222222222222222222222222222222222222222222222222222222"
|
||||
a := &Attachment{Hash: hash, OriginalFilename: "del.txt", Size: 10, MIMEType: "text/plain", UploadedBy: "agent-a"}
|
||||
if err := store.InsertMetadata(ctx, a); err != nil {
|
||||
t.Fatalf("InsertMetadata: %v", err)
|
||||
}
|
||||
|
||||
if err := store.DeleteByHash(ctx, hash); err != nil {
|
||||
t.Fatalf("DeleteByHash: %v", err)
|
||||
}
|
||||
|
||||
results, err := store.GetByHash(ctx, hash)
|
||||
if err != nil {
|
||||
t.Fatalf("GetByHash after delete: %v", err)
|
||||
}
|
||||
if len(results) != 0 {
|
||||
t.Errorf("expected 0 results after delete, got %d", len(results))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_CountReferences(t *testing.T) {
|
||||
store, _ := newTestStore(t)
|
||||
ctx := context.Background()
|
||||
|
||||
hash := "3333333333333333333333333333333333333333333333333333333333333333"
|
||||
|
||||
count, err := store.CountReferences(ctx, hash)
|
||||
if err != nil {
|
||||
t.Fatalf("CountReferences: %v", err)
|
||||
}
|
||||
if count != 0 {
|
||||
t.Errorf("expected 0 references, got %d", count)
|
||||
}
|
||||
|
||||
a := &Attachment{Hash: hash, OriginalFilename: "ref.txt", Size: 5, MIMEType: "text/plain", UploadedBy: "agent-a"}
|
||||
if err := store.InsertMetadata(ctx, a); err != nil {
|
||||
t.Fatalf("InsertMetadata: %v", err)
|
||||
}
|
||||
|
||||
count, err = store.CountReferences(ctx, hash)
|
||||
if err != nil {
|
||||
t.Fatalf("CountReferences: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("expected 1 reference, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteStore_FindOrphanHashes(t *testing.T) {
|
||||
store, db := newTestStore(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// Create an orphan attachment (no message_id).
|
||||
orphanHash := "4444444444444444444444444444444444444444444444444444444444444444"
|
||||
a := &Attachment{Hash: orphanHash, OriginalFilename: "orphan.txt", Size: 5, MIMEType: "text/plain", UploadedBy: "agent-a"}
|
||||
if err := store.InsertMetadata(ctx, a); err != nil {
|
||||
t.Fatalf("InsertMetadata orphan: %v", err)
|
||||
}
|
||||
|
||||
// Create a linked attachment.
|
||||
msgID := seedMessage(t, db)
|
||||
linkedHash := "5555555555555555555555555555555555555555555555555555555555555555"
|
||||
aLinked := &Attachment{Hash: linkedHash, OriginalFilename: "linked.txt", Size: 5, MIMEType: "text/plain", MessageID: &msgID, UploadedBy: "agent-a"}
|
||||
if err := store.InsertMetadata(ctx, aLinked); err != nil {
|
||||
t.Fatalf("InsertMetadata linked: %v", err)
|
||||
}
|
||||
|
||||
orphans, err := store.FindOrphanHashes(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("FindOrphanHashes: %v", err)
|
||||
}
|
||||
|
||||
if len(orphans) != 1 {
|
||||
t.Fatalf("expected 1 orphan, got %d", len(orphans))
|
||||
}
|
||||
if orphans[0] != orphanHash {
|
||||
t.Errorf("orphan hash = %s, want %s", orphans[0], orphanHash)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
package attachments
|
||||
|
||||
import "context"
|
||||
|
||||
// Store is the metadata persistence interface for attachments.
|
||||
type Store interface {
|
||||
// InsertMetadata inserts a new attachment metadata row.
|
||||
InsertMetadata(ctx context.Context, a *Attachment) error
|
||||
|
||||
// GetByHash returns all attachment metadata rows matching the given hash.
|
||||
GetByHash(ctx context.Context, hash string) ([]*Attachment, error)
|
||||
|
||||
// GetByMessageID returns all attachment metadata rows for a given message.
|
||||
GetByMessageID(ctx context.Context, messageID int64) ([]*Attachment, error)
|
||||
|
||||
// DeleteByHash removes all metadata rows matching the given hash.
|
||||
DeleteByHash(ctx context.Context, hash string) error
|
||||
|
||||
// FindOrphanHashes returns hashes that have no message_id reference
|
||||
// (message_id IS NULL or the referenced message no longer exists).
|
||||
FindOrphanHashes(ctx context.Context) ([]string, error)
|
||||
|
||||
// CountReferences returns the number of metadata rows referencing the hash.
|
||||
CountReferences(ctx context.Context, hash string) (int64, error)
|
||||
}
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/agents"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/attachments"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/channels"
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
|
||||
)
|
||||
@@ -26,6 +27,7 @@ func NewMCPServer(
|
||||
msgService *messaging.MessagingService,
|
||||
agentService *agents.AgentService,
|
||||
channelService *channels.Service,
|
||||
attachmentService *attachments.Service,
|
||||
) *MCPServer {
|
||||
logger := slog.Default().With("component", "mcp-server")
|
||||
|
||||
@@ -46,6 +48,12 @@ func NewMCPServer(
|
||||
channelRegistrar.RegisterAll(mcpSrv)
|
||||
}
|
||||
|
||||
// Register attachment tools
|
||||
if attachmentService != nil {
|
||||
attachmentRegistrar := NewAttachmentToolRegistrar(attachmentService)
|
||||
attachmentRegistrar.RegisterAll(mcpSrv)
|
||||
}
|
||||
|
||||
// Create SSE transport with context func for auth propagation
|
||||
sseServer := server.NewSSEServer(mcpSrv,
|
||||
server.WithSSEContextFunc(func(ctx context.Context, r *http.Request) context.Context {
|
||||
|
||||
@@ -0,0 +1,161 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/smart-mcp-proxy/synapbus/internal/attachments"
|
||||
)
|
||||
|
||||
// AttachmentToolRegistrar registers attachment MCP tools on the server.
|
||||
type AttachmentToolRegistrar struct {
|
||||
attachmentService *attachments.Service
|
||||
}
|
||||
|
||||
// NewAttachmentToolRegistrar creates a new attachment tool registrar.
|
||||
func NewAttachmentToolRegistrar(attachmentService *attachments.Service) *AttachmentToolRegistrar {
|
||||
return &AttachmentToolRegistrar{
|
||||
attachmentService: attachmentService,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterAll registers all attachment tools on the MCP server.
|
||||
func (atr *AttachmentToolRegistrar) RegisterAll(s *server.MCPServer) {
|
||||
s.AddTool(atr.uploadAttachmentTool(), atr.handleUploadAttachment)
|
||||
s.AddTool(atr.downloadAttachmentTool(), atr.handleDownloadAttachment)
|
||||
s.AddTool(atr.gcAttachmentsTool(), atr.handleGCAttachments)
|
||||
}
|
||||
|
||||
// --- Tool Definitions ---
|
||||
|
||||
func (atr *AttachmentToolRegistrar) uploadAttachmentTool() mcp.Tool {
|
||||
return mcp.NewTool("upload_attachment",
|
||||
mcp.WithDescription("Upload a file attachment. Content must be base64-encoded. Returns the SHA-256 hash for later retrieval. Max file size: 50MB."),
|
||||
mcp.WithString("content", mcp.Description("Base64-encoded file content"), mcp.Required()),
|
||||
mcp.WithString("filename", mcp.Description("Original filename (optional, used for MIME detection and display)")),
|
||||
mcp.WithString("mime_type", mcp.Description("MIME type override (optional, auto-detected from content if not provided)")),
|
||||
mcp.WithNumber("message_id", mcp.Description("Message ID to attach the file to (optional, can be linked later)")),
|
||||
)
|
||||
}
|
||||
|
||||
func (atr *AttachmentToolRegistrar) downloadAttachmentTool() mcp.Tool {
|
||||
return mcp.NewTool("download_attachment",
|
||||
mcp.WithDescription("Download an attachment by its SHA-256 hash. Returns base64-encoded content along with filename and MIME type metadata."),
|
||||
mcp.WithString("hash", mcp.Description("SHA-256 hash of the attachment"), mcp.Required()),
|
||||
)
|
||||
}
|
||||
|
||||
func (atr *AttachmentToolRegistrar) gcAttachmentsTool() mcp.Tool {
|
||||
return mcp.NewTool("gc_attachments",
|
||||
mcp.WithDescription("Run garbage collection to remove orphaned attachments not referenced by any message. Returns a summary of files removed and bytes reclaimed."),
|
||||
)
|
||||
}
|
||||
|
||||
// --- Tool Handlers ---
|
||||
|
||||
func (atr *AttachmentToolRegistrar) handleUploadAttachment(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
agentName, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
contentB64 := req.GetString("content", "")
|
||||
if contentB64 == "" {
|
||||
return mcp.NewToolResultError("'content' parameter is required"), nil
|
||||
}
|
||||
|
||||
// Check base64 size before decoding to avoid buffering oversized content.
|
||||
// Base64 expands data by ~4/3, so decoded size is roughly 3/4 of encoded.
|
||||
if int64(len(contentB64))*3/4 > attachments.MaxFileSize {
|
||||
return mcp.NewToolResultError("file exceeds maximum size of 50MB"), nil
|
||||
}
|
||||
|
||||
decoded, err := base64.StdEncoding.DecodeString(contentB64)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("invalid base64 content: %s", err)), nil
|
||||
}
|
||||
|
||||
if int64(len(decoded)) > attachments.MaxFileSize {
|
||||
return mcp.NewToolResultError("file exceeds maximum size of 50MB"), nil
|
||||
}
|
||||
|
||||
uploadReq := attachments.UploadRequest{
|
||||
Content: bytes.NewReader(decoded),
|
||||
Filename: req.GetString("filename", ""),
|
||||
MIMEType: req.GetString("mime_type", ""),
|
||||
UploadedBy: agentName,
|
||||
}
|
||||
|
||||
// Optional message_id.
|
||||
if mid := req.GetInt("message_id", 0); mid > 0 {
|
||||
v := int64(mid)
|
||||
uploadReq.MessageID = &v
|
||||
}
|
||||
|
||||
result, err := atr.attachmentService.Upload(ctx, uploadReq)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("upload_attachment failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"hash": result.Hash,
|
||||
"size": result.Size,
|
||||
"mime_type": result.MIMEType,
|
||||
"original_filename": result.Filename,
|
||||
})
|
||||
}
|
||||
|
||||
func (atr *AttachmentToolRegistrar) handleDownloadAttachment(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
_, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
hash := req.GetString("hash", "")
|
||||
if hash == "" {
|
||||
return mcp.NewToolResultError("'hash' parameter is required"), nil
|
||||
}
|
||||
|
||||
result, err := atr.attachmentService.Download(ctx, hash)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("download_attachment failed: %s", err)), nil
|
||||
}
|
||||
defer result.Content.Close()
|
||||
|
||||
// Read content and base64-encode it.
|
||||
content, err := io.ReadAll(result.Content)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("read attachment content failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"hash": result.Hash,
|
||||
"content": base64.StdEncoding.EncodeToString(content),
|
||||
"original_filename": result.Filename,
|
||||
"mime_type": result.MIMEType,
|
||||
"size": result.Size,
|
||||
})
|
||||
}
|
||||
|
||||
func (atr *AttachmentToolRegistrar) handleGCAttachments(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
_, ok := extractAgentName(ctx)
|
||||
if !ok {
|
||||
return mcp.NewToolResultError("authentication required"), nil
|
||||
}
|
||||
|
||||
result, err := atr.attachmentService.GarbageCollect(ctx)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("gc_attachments failed: %s", err)), nil
|
||||
}
|
||||
|
||||
return resultJSON(map[string]any{
|
||||
"files_removed": result.FilesRemoved,
|
||||
"bytes_reclaimed": result.BytesReclaimed,
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user