260 lines
8.0 KiB
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
|
|
}
|