Changeset 0.13.1 (#65)
This commit is contained in:
308
server/tools/urlfetch.go
Normal file
308
server/tools/urlfetch.go
Normal file
@@ -0,0 +1,308 @@
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"golang.org/x/net/html"
|
||||
)
|
||||
|
||||
func init() {
|
||||
Register(&URLFetchTool{})
|
||||
}
|
||||
|
||||
// ═══════════════════════════════════════════
|
||||
// url_fetch
|
||||
// ═══════════════════════════════════════════
|
||||
|
||||
type URLFetchTool struct{}
|
||||
|
||||
func (t *URLFetchTool) Definition() ToolDef {
|
||||
return ToolDef{
|
||||
Name: "url_fetch",
|
||||
DisplayName: "Fetch URL",
|
||||
Description: "Fetch and extract readable text content from a web page URL. Use this to read articles, documentation, or any web page the user shares or that appeared in search results. Returns the page title and extracted text.",
|
||||
Category: "search",
|
||||
Parameters: JSONSchema(map[string]interface{}{
|
||||
"url": Prop("string", "Full URL to fetch (must start with http:// or https://)"),
|
||||
"max_length": Prop("integer", "Maximum content length in characters (default 8000, max 16000)"),
|
||||
}, []string{"url"}),
|
||||
}
|
||||
}
|
||||
|
||||
func (t *URLFetchTool) Execute(ctx context.Context, execCtx ExecutionContext, argsJSON string) (string, error) {
|
||||
var args struct {
|
||||
URL string `json:"url"`
|
||||
MaxLength int `json:"max_length"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(argsJSON), &args); err != nil {
|
||||
return "", fmt.Errorf("invalid arguments: %w", err)
|
||||
}
|
||||
|
||||
if args.URL == "" {
|
||||
return "", fmt.Errorf("url is required")
|
||||
}
|
||||
if !strings.HasPrefix(args.URL, "http://") && !strings.HasPrefix(args.URL, "https://") {
|
||||
return "", fmt.Errorf("url must start with http:// or https://")
|
||||
}
|
||||
|
||||
if args.MaxLength <= 0 {
|
||||
args.MaxLength = 8000
|
||||
}
|
||||
if args.MaxLength > 16000 {
|
||||
args.MaxLength = 16000
|
||||
}
|
||||
|
||||
// SSRF protection: resolve hostname and block private IPs
|
||||
if err := checkSSRF(args.URL); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// Fetch with timeout
|
||||
client := &http.Client{
|
||||
Timeout: 15 * time.Second,
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
if len(via) >= 3 {
|
||||
return fmt.Errorf("too many redirects")
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", args.URL, nil)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("invalid URL: %w", err)
|
||||
}
|
||||
req.Header.Set("User-Agent", "ChatSwitchboard/1.0 (URL Fetch Tool)")
|
||||
req.Header.Set("Accept", "text/html,application/xhtml+xml,text/plain")
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("fetch failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("HTTP %d: %s", resp.StatusCode, resp.Status)
|
||||
}
|
||||
|
||||
ct := resp.Header.Get("Content-Type")
|
||||
if !strings.Contains(ct, "text/html") && !strings.Contains(ct, "text/plain") &&
|
||||
!strings.Contains(ct, "application/xhtml") && !strings.Contains(ct, "application/json") {
|
||||
return "", fmt.Errorf("unsupported content type: %s (expected HTML or text)", ct)
|
||||
}
|
||||
|
||||
// Read body with limit (1MB raw)
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, 1024*1024))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("read body: %w", err)
|
||||
}
|
||||
|
||||
// Plain text or JSON → return as-is (truncated)
|
||||
if strings.Contains(ct, "text/plain") || strings.Contains(ct, "application/json") {
|
||||
content := string(body)
|
||||
content = truncateUTF8(content, args.MaxLength)
|
||||
result := map[string]interface{}{
|
||||
"url": args.URL,
|
||||
"content_type": ct,
|
||||
"content": content,
|
||||
"truncated": len(body) > args.MaxLength,
|
||||
}
|
||||
b, _ := json.Marshal(result)
|
||||
return string(b), nil
|
||||
}
|
||||
|
||||
// HTML → extract readable text
|
||||
title, text := extractReadableHTML(string(body))
|
||||
text = truncateUTF8(text, args.MaxLength)
|
||||
|
||||
result := map[string]interface{}{
|
||||
"url": args.URL,
|
||||
"title": title,
|
||||
"content": text,
|
||||
"truncated": utf8.RuneCountInString(string(body)) > args.MaxLength,
|
||||
}
|
||||
|
||||
b, _ := json.Marshal(result)
|
||||
return string(b), nil
|
||||
}
|
||||
|
||||
// ── SSRF Protection ─────────────────────────
|
||||
|
||||
func checkSSRF(rawURL string) error {
|
||||
// Parse hostname
|
||||
host := rawURL
|
||||
if idx := strings.Index(rawURL, "://"); idx >= 0 {
|
||||
host = rawURL[idx+3:]
|
||||
}
|
||||
if idx := strings.IndexAny(host, ":/"); idx >= 0 {
|
||||
host = host[:idx]
|
||||
}
|
||||
|
||||
// Resolve DNS
|
||||
ips, err := net.LookupIP(host)
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot resolve host: %s", host)
|
||||
}
|
||||
|
||||
for _, ip := range ips {
|
||||
if isPrivateIP(ip) {
|
||||
return fmt.Errorf("blocked: %s resolves to private IP %s", host, ip.String())
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func isPrivateIP(ip net.IP) bool {
|
||||
privateRanges := []struct {
|
||||
network *net.IPNet
|
||||
}{
|
||||
{mustParseCIDR("10.0.0.0/8")},
|
||||
{mustParseCIDR("172.16.0.0/12")},
|
||||
{mustParseCIDR("192.168.0.0/16")},
|
||||
{mustParseCIDR("127.0.0.0/8")},
|
||||
{mustParseCIDR("169.254.0.0/16")},
|
||||
{mustParseCIDR("::1/128")},
|
||||
{mustParseCIDR("fc00::/7")},
|
||||
{mustParseCIDR("fe80::/10")},
|
||||
}
|
||||
for _, r := range privateRanges {
|
||||
if r.network.Contains(ip) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func mustParseCIDR(s string) *net.IPNet {
|
||||
_, net, err := net.ParseCIDR(s)
|
||||
if err != nil {
|
||||
panic("invalid CIDR: " + s)
|
||||
}
|
||||
return net
|
||||
}
|
||||
|
||||
// ── HTML → Text Extraction ──────────────────
|
||||
|
||||
// extractReadableHTML extracts the page title and readable text from HTML.
|
||||
// Strips scripts, styles, nav, header, footer elements.
|
||||
func extractReadableHTML(htmlContent string) (title, text string) {
|
||||
doc, err := html.Parse(strings.NewReader(htmlContent))
|
||||
if err != nil {
|
||||
return "", htmlContent
|
||||
}
|
||||
|
||||
// Extract <title>
|
||||
title = findTitle(doc)
|
||||
|
||||
// Extract text from body, skipping non-content elements
|
||||
var sb strings.Builder
|
||||
extractText(doc, &sb, false)
|
||||
text = cleanText(sb.String())
|
||||
|
||||
return title, text
|
||||
}
|
||||
|
||||
func findTitle(n *html.Node) string {
|
||||
if n.Type == html.ElementNode && n.Data == "title" {
|
||||
return strings.TrimSpace(textContentNode(n))
|
||||
}
|
||||
for c := n.FirstChild; c != nil; c = c.NextSibling {
|
||||
if t := findTitle(c); t != "" {
|
||||
return t
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// skipTags are elements whose content is not readable text.
|
||||
var skipTags = map[string]bool{
|
||||
"script": true, "style": true, "noscript": true,
|
||||
"nav": true, "header": true, "footer": true,
|
||||
"aside": true, "svg": true, "form": true,
|
||||
"iframe": true, "button": true, "select": true,
|
||||
}
|
||||
|
||||
func extractText(n *html.Node, sb *strings.Builder, inBody bool) {
|
||||
if n.Type == html.ElementNode {
|
||||
if skipTags[n.Data] {
|
||||
return
|
||||
}
|
||||
if n.Data == "body" {
|
||||
inBody = true
|
||||
}
|
||||
}
|
||||
|
||||
if inBody && n.Type == html.TextNode {
|
||||
text := strings.TrimSpace(n.Data)
|
||||
if text != "" {
|
||||
sb.WriteString(text)
|
||||
sb.WriteString(" ")
|
||||
}
|
||||
}
|
||||
|
||||
// Block elements get a newline
|
||||
if n.Type == html.ElementNode {
|
||||
switch n.Data {
|
||||
case "p", "div", "br", "h1", "h2", "h3", "h4", "h5", "h6",
|
||||
"li", "tr", "blockquote", "pre", "article", "section":
|
||||
sb.WriteString("\n")
|
||||
}
|
||||
}
|
||||
|
||||
for c := n.FirstChild; c != nil; c = c.NextSibling {
|
||||
extractText(c, sb, inBody)
|
||||
}
|
||||
}
|
||||
|
||||
func textContentNode(n *html.Node) string {
|
||||
if n.Type == html.TextNode {
|
||||
return n.Data
|
||||
}
|
||||
var sb strings.Builder
|
||||
for c := n.FirstChild; c != nil; c = c.NextSibling {
|
||||
sb.WriteString(textContentNode(c))
|
||||
}
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
// cleanText collapses multiple spaces/newlines into single separators.
|
||||
func cleanText(s string) string {
|
||||
var sb strings.Builder
|
||||
prevNewline := false
|
||||
prevSpace := false
|
||||
for _, r := range s {
|
||||
if r == '\n' {
|
||||
if !prevNewline {
|
||||
sb.WriteRune('\n')
|
||||
}
|
||||
prevNewline = true
|
||||
prevSpace = false
|
||||
} else if r == ' ' || r == '\t' {
|
||||
if !prevSpace && !prevNewline {
|
||||
sb.WriteRune(' ')
|
||||
}
|
||||
prevSpace = true
|
||||
} else {
|
||||
sb.WriteRune(r)
|
||||
prevNewline = false
|
||||
prevSpace = false
|
||||
}
|
||||
}
|
||||
return strings.TrimSpace(sb.String())
|
||||
}
|
||||
|
||||
func truncateUTF8(s string, maxChars int) string {
|
||||
runes := []rune(s)
|
||||
if len(runes) <= maxChars {
|
||||
return s
|
||||
}
|
||||
return string(runes[:maxChars]) + "\n[content truncated]"
|
||||
}
|
||||
Reference in New Issue
Block a user