Files
ombi-mcp/internal/tools/call.go

260 lines
8.0 KiB
Go

package tools
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
"strconv"
"strings"
"ombi-mcp/internal/ombi"
)
// maxUpstreamBody bounds how much of one upstream response is read
// before projection; oversized data is refused rather than parsed
// into an unbounded entity graph.
const maxUpstreamBody = 8 << 20 // 8 MiB
// maxSanitizedMsg bounds upstream-derived error text.
const maxSanitizedMsg = 300
// seg percent-encodes one path segment value.
func seg(v string) string { return url.PathEscape(v) }
func segInt(v int) string { return strconv.Itoa(v) }
// args strictly decodes the MCP arguments object into v. Unknown
// properties are rejected (additionalProperties:false semantics).
// Returns a finished error envelope, or nil when decoding succeeds.
func (o *op) args(raw json.RawMessage, v any) *ToolResult {
if len(raw) == 0 {
return o.invalid("", "arguments object is required")
}
dec := json.NewDecoder(bytes.NewReader(raw))
dec.DisallowUnknownFields()
if err := dec.Decode(v); err != nil {
return o.invalid("", "invalid arguments: %s", sanitizeErr(err))
}
return nil
}
// call performs one upstream request. On HTTP 2xx it returns the raw
// body; otherwise it returns a finished error envelope. The caller
// must return the envelope immediately when non-nil.
func (o *op) call(method, path string, query map[string]string, body any) ([]byte, *ToolResult) {
resp, err := o.env.Upstream.Do(o.ctx, method, path, query, body)
if err != nil {
return nil, o.transportErr(err)
}
defer resp.Body.Close()
raw, err := io.ReadAll(io.LimitReader(resp.Body, maxUpstreamBody+1))
if err != nil {
return nil, o.transportErr(err)
}
if len(raw) > maxUpstreamBody {
return nil, o.fail("UPSTREAM_SCHEMA_MISMATCH",
"upstream response exceeded the read budget", false)
}
if resp.StatusCode < 200 || resp.StatusCode > 299 {
return nil, o.httpErr(resp, raw)
}
return raw, nil
}
// transportErr maps client/transport failures onto ToolError codes.
// A timed-out or disconnected write has unknown outcome and is never
// marked retryable.
func (o *op) transportErr(err error) *ToolResult {
switch {
case errors.Is(err, ombi.ErrUnauthorized), errors.Is(err, ombi.ErrAuthFailed):
st := http.StatusUnauthorized
return o.failErr(&ToolError{
Code: "AUTHENTICATION_FAILED",
Message: "upstream authentication failed",
Retryable: false,
HTTPStatus: &st,
})
case errors.Is(err, context.DeadlineExceeded) || isNetTimeout(err):
if o.write {
return o.fail("UNKNOWN_OUTCOME",
"upstream timed out after the request may have executed", false)
}
return o.fail("TIMEOUT", "upstream request timed out", true)
default:
if o.write {
return o.fail("UNKNOWN_OUTCOME",
"upstream transport failure; outcome is unknown", false)
}
return o.fail("UPSTREAM_REJECTED",
fmt.Sprintf("upstream transport failure: %s", sanitizeErr(err)), true)
}
}
func isNetTimeout(err error) bool {
var ne net.Error
return errors.As(err, &ne) && ne.Timeout()
}
// httpErr maps an upstream non-2xx status to a sanitized ToolError.
// Raw bodies are never forwarded; a short sanitized message is kept
// only when the body parses as JSON.
func (o *op) httpErr(resp *http.Response, raw []byte) *ToolResult {
st := resp.StatusCode
e := &ToolError{HTTPStatus: &st}
switch st {
case http.StatusBadRequest, http.StatusUnprocessableEntity:
e.Code, e.Message, e.Retryable = "UPSTREAM_REJECTED",
"upstream rejected the request"+sanitizedDetail(raw), false
case http.StatusUnauthorized:
e.Code, e.Message, e.Retryable = "AUTHENTICATION_FAILED",
"upstream authentication failed", false
case http.StatusForbidden:
e.Code, e.Message, e.Retryable = "FORBIDDEN",
"upstream denied the request", false
case http.StatusNotFound:
e.Code, e.Message, e.Retryable = "NOT_FOUND",
"upstream resource not found"+sanitizedDetail(raw), false
case http.StatusConflict:
e.Code, e.Message, e.Retryable = "CONFLICT",
"upstream conflict"+sanitizedDetail(raw), false
case http.StatusTooManyRequests:
e.Code, e.Message, e.Retryable = "RATE_LIMITED",
"upstream rate limit exceeded", true
if s := parseRetryAfter(resp.Header.Get("Retry-After")); s != nil {
e.RetryAfterSeconds = s
}
default:
if st >= 500 {
e.Code, e.Message, e.Retryable = "UPSTREAM_REJECTED",
fmt.Sprintf("upstream error (HTTP %d)", st)+sanitizedDetail(raw), true
} else {
e.Code, e.Message, e.Retryable = "UPSTREAM_REJECTED",
fmt.Sprintf("upstream returned HTTP %d", st), false
}
}
return o.failErr(e)
}
// sanitizedDetail extracts a short human message from an upstream
// error body. It returns "" for empty/HTML/oversized/unparseable
// bodies — never forwarding stack traces, HTML pages or raw payloads.
func sanitizedDetail(raw []byte) string {
trim := bytes.TrimSpace(raw)
if len(trim) == 0 || trim[0] == '<' {
return ""
}
var m map[string]any
if err := json.Unmarshal(trim, &m); err != nil {
return ""
}
for _, k := range []string{"title", "errorMessage", "message", "Message", "detail", "error"} {
if s, ok := m[k].(string); ok {
if s = sanitizeText(s, maxSanitizedMsg); s != "" {
return ": " + s
}
}
}
return ""
}
// sanitizeText trims a string to n chars, strips control characters
// and collapses whitespace so nothing multiline or HTML-like leaks.
func sanitizeText(s string, n int) string {
s = strings.Map(func(r rune) rune {
if r < 0x20 || r == 0x7f {
return ' '
}
return r
}, s)
s = strings.Join(strings.Fields(s), " ")
if len(s) > n {
s = s[:n] + "…"
}
return s
}
// sanitizeErr renders a Go error without exposing internals beyond a
// bounded single-line message.
func sanitizeErr(err error) string { return sanitizeText(err.Error(), maxSanitizedMsg) }
func parseRetryAfter(v string) *int {
if v == "" {
return nil
}
if n, err := strconv.Atoi(strings.TrimSpace(v)); err == nil && n >= 0 {
return &n
}
return nil
}
// decodeJSON unmarshals an upstream body into v. Decode failures map
// to UPSTREAM_SCHEMA_MISMATCH rather than fabricated defaults.
func (o *op) decodeJSON(raw []byte, v any) *ToolResult {
if err := json.Unmarshal(raw, v); err != nil {
return o.fail("UPSTREAM_SCHEMA_MISMATCH",
fmt.Sprintf("upstream response did not match the expected shape: %s", sanitizeErr(err)),
false)
}
return nil
}
// decodeObject unmarshals an upstream body expecting a JSON object.
func (o *op) decodeObject(raw []byte) (map[string]any, *ToolResult) {
var m map[string]any
if err := json.Unmarshal(raw, &m); err != nil {
return nil, o.fail("UPSTREAM_SCHEMA_MISMATCH",
fmt.Sprintf("upstream response was not a JSON object: %s", sanitizeErr(err)), false)
}
if m == nil {
return nil, o.fail("UPSTREAM_SCHEMA_MISMATCH", "upstream response was null", false)
}
return m, nil
}
// decodeArray unmarshals an upstream body expecting a JSON array.
func (o *op) decodeArray(raw []byte) ([]map[string]any, *ToolResult) {
var arr []map[string]any
if err := json.Unmarshal(raw, &arr); err != nil {
return nil, o.fail("UPSTREAM_SCHEMA_MISMATCH",
fmt.Sprintf("upstream response was not a JSON array: %s", sanitizeErr(err)), false)
}
if arr == nil {
arr = []map[string]any{}
}
return arr, nil
}
// decodeScalar unmarshals an upstream body that is a bare JSON scalar
// (bool, number or string). Unquoted plain text is tolerated as a
// string for upstream routes that return raw text.
func (o *op) decodeScalar(raw []byte) (any, *ToolResult) {
trim := bytes.TrimSpace(raw)
if len(trim) == 0 {
return nil, nil
}
var v any
if err := json.Unmarshal(trim, &v); err == nil {
return v, nil
}
if len(trim) > 1<<20 {
return nil, o.fail("UPSTREAM_SCHEMA_MISMATCH", "upstream scalar body too large", false)
}
return string(trim), nil
}
// decodeBool unmarshals an upstream body expecting a JSON boolean.
func (o *op) decodeBool(raw []byte) (bool, *ToolResult) {
var b bool
if err := json.Unmarshal(bytes.TrimSpace(raw), &b); err != nil {
return false, o.fail("UPSTREAM_SCHEMA_MISMATCH",
"upstream response was not a boolean", false)
}
return b, nil
}