Files
ha-gronod-addons/emby-mcp/internal/bridge/bridge.go
gronod 9149ec844a fix: filter stale REST args against tool schema, keep lyrics_or_description (1.0.12)
The HA conversation agent still posts lyrics_or_description on every
search_for_item call; the go-sdk's additionalProperties:false validation
rejected the whole request, breaking search. The REST bridge now caches
each tool's input schema from tools/list and drops undeclared arguments,
reporting them as dropped_arguments, while search_for_item re-accepts the
param as a deprecated no-op for direct MCP callers.
2026-09-22 00:24:50 +01:00

412 lines
12 KiB
Go

// Package bridge exposes a small REST API that Home Assistant (and other
// clients) can call to invoke Emby.MCP tools without speaking MCP directly.
//
// Credentials are forwarded from each incoming request. The bridge does not
// cache Emby access tokens; the MCP authenticator already does that.
package bridge
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"sort"
"strings"
"sync"
"time"
"git.i3omb.com/gronod/emby-mcp/internal/applog"
"git.i3omb.com/gronod/emby-mcp/internal/config"
"git.i3omb.com/gronod/emby-mcp/internal/mcphttp"
)
// Handler serves the REST bridge endpoints under /call/{tool} and /tools.
type Handler struct {
cfg *config.Config
client *http.Client
schemaMu sync.Mutex
schemas map[string]map[string]bool // tool name → declared argument names
}
// NewHandler builds a bridge handler.
func NewHandler(cfg *config.Config) *Handler {
return &Handler{
cfg: cfg,
client: &http.Client{
Timeout: 30 * time.Second,
Transport: &http.Transport{
MaxIdleConns: 16,
MaxIdleConnsPerHost: 8,
IdleConnTimeout: 90 * time.Second,
},
},
}
}
// Register mounts the bridge routes onto mux.
func (h *Handler) Register(mux *http.ServeMux) {
mux.Handle("GET /tools", restLog(http.HandlerFunc(h.listTools)))
mux.Handle("POST /call/{tool}", restLog(http.HandlerFunc(h.callTool)))
mux.HandleFunc("GET /health", func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"status":"ok"}`))
})
}
type session struct {
id string
initAt time.Time
}
var (
sessMu sync.Mutex
sess *session
)
func (h *Handler) resetSession() {
sessMu.Lock()
sess = nil
sessMu.Unlock()
}
func (h *Handler) ensureSession(ctx context.Context, auth string) (string, error) {
sessMu.Lock()
defer sessMu.Unlock()
timeout := h.cfg.SessionTimeout
if timeout <= 0 {
timeout = 30 * time.Minute
}
if sess != nil && time.Since(sess.initAt) < timeout {
return sess.id, nil
}
id, err := h.initialize(ctx, auth)
if err != nil {
sess = nil
return "", err
}
sess = &session{id: id, initAt: time.Now()}
return id, nil
}
func (h *Handler) initialize(ctx context.Context, auth string) (string, error) {
payload := map[string]any{
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": map[string]any{
"protocolVersion": "2025-03-26",
"capabilities": map[string]any{},
"clientInfo": map[string]any{"name": "ha-emby-bridge", "version": "1.0.0"},
},
}
body, _ := json.Marshal(payload)
req, err := http.NewRequestWithContext(ctx, http.MethodPost, mcphttp.EndpointURL(), bytes.NewReader(body))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json, text/event-stream")
if auth != "" {
req.Header.Set("Authorization", auth)
}
resp, err := h.client.Do(req)
if err != nil {
return "", fmt.Errorf("initialize: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
b, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
return "", fmt.Errorf("initialize: status %d: %s", resp.StatusCode, strings.TrimSpace(string(b)))
}
id := resp.Header.Get("Mcp-Session-Id")
if id == "" {
return "", fmt.Errorf("initialize: missing Mcp-Session-Id header")
}
notif, _ := json.Marshal(map[string]any{"jsonrpc": "2.0", "method": "notifications/initialized"})
nreq, err := http.NewRequestWithContext(ctx, http.MethodPost, mcphttp.EndpointURL(), bytes.NewReader(notif))
if err == nil {
nreq.Header.Set("Content-Type", "application/json")
nreq.Header.Set("Accept", "application/json, text/event-stream")
nreq.Header.Set("Mcp-Session-Id", id)
if auth != "" {
nreq.Header.Set("Authorization", auth)
}
if nr, err := h.client.Do(nreq); err == nil {
nr.Body.Close()
}
}
return id, nil
}
func (h *Handler) post(ctx context.Context, sid, auth string, payload any) (*http.Response, error) {
body, err := json.Marshal(payload)
if err != nil {
return nil, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, mcphttp.EndpointURL(), bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json, text/event-stream")
if sid != "" {
req.Header.Set("Mcp-Session-Id", sid)
}
if auth != "" {
req.Header.Set("Authorization", auth)
}
return h.client.Do(req)
}
func retryable(code int) bool {
return code == http.StatusBadRequest || code == http.StatusUnauthorized || code == http.StatusNotFound
}
func (h *Handler) listTools(w http.ResponseWriter, r *http.Request) {
h.rpc(w, r, map[string]any{"jsonrpc": "2.0", "method": "tools/list", "id": 2}, func(data map[string]any) {
tools := []any{}
if res, ok := data["result"].(map[string]any); ok {
if t, ok := res["tools"].([]any); ok {
tools = t
}
}
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(tools)
})
}
func (h *Handler) callTool(w http.ResponseWriter, r *http.Request) {
tool := r.PathValue("tool")
var args map[string]any
_ = json.NewDecoder(r.Body).Decode(&args)
if args == nil {
args = map[string]any{}
}
dropped := h.filterArgs(r.Context(), r.Header.Get("Authorization"), tool, args)
payload := map[string]any{
"jsonrpc": "2.0",
"method": "tools/call",
"params": map[string]any{"name": tool, "arguments": args},
"id": 3,
}
h.rpc(w, r, payload, func(data map[string]any) {
out := "Success"
if res, ok := data["result"].(map[string]any); ok {
if content, ok := res["content"].([]any); ok {
var parts []string
for _, c := range content {
if m, ok := c.(map[string]any); ok && m["type"] == "text" {
if t, ok := m["text"].(string); ok && t != "" {
parts = append(parts, t)
}
}
}
if len(parts) > 0 {
out = strings.Join(parts, "\n")
}
}
}
resp := map[string]any{"result": out}
if len(dropped) > 0 {
resp["dropped_arguments"] = dropped
}
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(resp)
})
}
// toolArgs returns the declared argument names for tool, fetching and caching
// tools/list on first use. Returns nil when the schema is unavailable or the
// tool is unknown — callers then pass arguments through unfiltered and let
// MCP-side validation decide.
func (h *Handler) toolArgs(ctx context.Context, auth, tool string) map[string]bool {
h.schemaMu.Lock()
defer h.schemaMu.Unlock()
if h.schemas == nil {
data, err := h.rpcCall(ctx, auth,
map[string]any{"jsonrpc": "2.0", "method": "tools/list", "id": 2})
if err != nil {
applog.Warnf("bridge: tools/list failed, arguments unfiltered: %v", err)
return nil
}
h.schemas = map[string]map[string]bool{}
if res, ok := data["result"].(map[string]any); ok {
if tools, ok := res["tools"].([]any); ok {
for _, t := range tools {
m, _ := t.(map[string]any)
name, _ := m["name"].(string)
schema, _ := m["inputSchema"].(map[string]any)
props, _ := schema["properties"].(map[string]any)
allowed := make(map[string]bool, len(props))
for k := range props {
allowed[k] = true
}
if name != "" {
h.schemas[name] = allowed
}
}
}
}
}
return h.schemas[tool]
}
// filterArgs removes keys from args that the tool's input schema does not
// declare, returning the sorted list of dropped keys. Unavailable schemas
// leave args untouched.
func (h *Handler) filterArgs(ctx context.Context, auth, tool string, args map[string]any) []string {
if len(args) == 0 {
return nil
}
allowed := h.toolArgs(ctx, auth, tool)
if allowed == nil {
return nil
}
var dropped []string
for k := range args {
if !allowed[k] {
delete(args, k)
dropped = append(dropped, k)
}
}
if len(dropped) > 0 {
sort.Strings(dropped)
applog.Debugf("bridge: dropped undeclared arguments for %s: %s", tool, strings.Join(dropped, ","))
}
return dropped
}
// rpcCall performs a JSON-RPC call on the shared MCP session: it posts the
// payload, retries once with a fresh session on a retryable HTTP status, and
// retries once more when the JSON-RPC error looks like an auth failure.
// The returned map may still contain a JSON-RPC "error" member.
func (h *Handler) rpcCall(ctx context.Context, auth string, payload map[string]any) (map[string]any, error) {
data, err := h.postRPC(ctx, auth, payload)
if err != nil {
return nil, err
}
if e, ok := data["error"]; ok && isAuthError(e) {
h.resetSession()
data, err = h.postRPC(ctx, auth, payload)
}
return data, err
}
// postRPC ensures a session, posts payload, and parses the SSE response,
// retrying once with a fresh session on a retryable HTTP status.
func (h *Handler) postRPC(ctx context.Context, auth string, payload map[string]any) (map[string]any, error) {
sid, err := h.ensureSession(ctx, auth)
if err != nil {
return nil, fmt.Errorf("session init failed: %w", err)
}
resp, err := h.post(ctx, sid, auth, payload)
if err != nil {
return nil, fmt.Errorf("MCP error: %w", err)
}
if retryable(resp.StatusCode) {
resp.Body.Close()
h.resetSession()
sid, err = h.ensureSession(ctx, auth)
if err != nil {
return nil, fmt.Errorf("session re-init failed: %w", err)
}
resp, err = h.post(ctx, sid, auth, payload)
if err != nil {
return nil, fmt.Errorf("MCP error: %w", err)
}
}
defer resp.Body.Close()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
b, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
return nil, fmt.Errorf("MCP status %d: %s", resp.StatusCode, strings.TrimSpace(string(b)))
}
return parseSSE(resp.Body)
}
func (h *Handler) rpc(w http.ResponseWriter, r *http.Request, payload map[string]any, okFn func(map[string]any)) {
data, err := h.rpcCall(r.Context(), r.Header.Get("Authorization"), payload)
if err != nil {
http.Error(w, err.Error(), http.StatusBadGateway)
return
}
if e, ok := data["error"]; ok {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusBadRequest)
_ = json.NewEncoder(w).Encode(map[string]any{"error": e})
return
}
okFn(data)
}
func isAuthError(e any) bool {
text := strings.ToLower(fmt.Sprintf("%v", e))
for _, n := range []string{"401", "unauthorized", "token", "expired", "access token"} {
if strings.Contains(text, n) {
return true
}
}
return false
}
func parseSSE(r io.Reader) (map[string]any, error) {
b, err := io.ReadAll(r)
if err != nil {
return nil, err
}
text := strings.TrimSpace(string(b))
if strings.Contains(text, "data:") {
var lines []string
for _, ln := range strings.Split(text, "\n") {
ln = strings.TrimSpace(ln)
if strings.HasPrefix(ln, "data:") {
lines = append(lines, strings.TrimSpace(ln[5:]))
}
}
if len(lines) > 0 {
text = strings.Join(lines, "\n")
}
}
var out map[string]any
if err := json.Unmarshal([]byte(text), &out); err != nil {
raw := text
if len(raw) > 500 {
raw = raw[:500]
}
return nil, fmt.Errorf("%v (raw: %s)", err, raw)
}
return out, nil
}
func restLog(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var reqBody []byte
if r.Body != nil {
reqBody, _ = io.ReadAll(io.LimitReader(r.Body, 1<<20))
r.Body = io.NopCloser(bytes.NewReader(reqBody))
}
buf := &bytes.Buffer{}
cw := &restCapture{ResponseWriter: w, buf: buf, code: 200}
next.ServeHTTP(cw, r)
applog.REST("%s %s req=%s resp=%s", r.Method, r.URL.RequestURI(), applog.Body(reqBody), applog.Body(buf.Bytes()))
})
}
type restCapture struct {
http.ResponseWriter
buf *bytes.Buffer
code int
}
func (w *restCapture) Write(p []byte) (int, error) {
_, _ = w.buf.Write(p)
return w.ResponseWriter.Write(p)
}
func (w *restCapture) WriteHeader(code int) {
w.code = code
w.ResponseWriter.WriteHeader(code)
}