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.
412 lines
12 KiB
Go
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)
|
|
}
|