Skip to main content

yggdrasil/api/
url_fetch.rs

1//! SSRF 防护的 URL 抓取:服务端按图,供 MCP `upload_media(url)` 工具使用。
2//!
3//! 这是 Option B 的第三通道——LLM 自身(不经 host/shell)绕过 base64 的唯一方式:
4//! 工具只收 `url: String`(JSON-RPC 纯文本),服务端去抓二进制,二进制从不进 JSON-RPC。
5//!
6//! # SSRF 防护(多层纵深)
7//!
8//! 1. **仅 https**:拒绝 http:// 与其它 scheme(公共图床几乎都走 https;
9//!    限 scheme 是最低成本的最大攻击面收敛)。
10//! 2. **解析即锁定 IP**:先用 `std::net::ToSocketAddrs` 解析主机名,逐个 IP 做
11//!    私网/回环/链路本地/保留段校验,**只把通过校验的 IP** 经 reqwest 的
12//!    `.resolve(host, ip)` 钉死——reqwest 不再二次 DNS 查询,杜绝 DNS rebinding
13//!    (解析时返回公网 IP 骗过校验,连接瞬间返回内网 IP 的攻击)。
14//! 3. **禁重定向**:302 可重定向到内网 IP 绕过白名单——设 `redirect(Policy::none)`,
15//!    抓取器对每个跳转点零信任。
16//! 4. **体积上限**:流式读 + 累计字节,超过 `MAX_FILE_SIZE` 立即中止,避免无界
17//!    响应(Content-Length 可伪造、可缺省)撑爆内存。
18//! 5. **超时**:连接 + 整体读取封顶,防慢速 loris 式拖挂。
19//! 6. **白名单兜底**:最终交给 `process_image_upload`,magic bytes 二次验真,
20//!    Content-Type 不被信任。
21//!
22//! 抓取在异步 reqwest 中执行(不阻塞 worker;转码/校验仍走 spawn_blocking)。
23//! 仅 `feature = "server"` 编译。
24
25#![cfg(feature = "server")]
26
27use std::net::{IpAddr, SocketAddr, ToSocketAddrs};
28use std::time::Duration;
29
30use bytes::Bytes;
31use http::Uri;
32
33use crate::api::upload::{process_image_upload, UploadError, UploadOutcome};
34use crate::utils::server::MAX_FILE_SIZE;
35
36/// URL 抓取错误(脱敏,映射到 MCP invalid_request / internal_error)。
37pub(crate) enum FetchError {
38    /// 400 类:scheme 非 https / 主机缺失 / 无法解析 / 指向禁止 IP。
39    Invalid(String),
40    /// 400 类:远端返回非成功状态、Content-Type 非图片。
41    BadStatus(String),
42    /// 413 类:响应体超过上限。
43    TooLarge,
44    /// 500 类:网络/IO/超时等(日志记详情,客户端见静态消息)。
45    Fetch(&'static str),
46}
47
48impl FetchError {
49    fn invalid<E: std::fmt::Display>(e: E, ctx: &str) -> Self {
50        tracing::warn!("url_fetch {ctx}: {e}");
51        FetchError::Invalid(format!("无效的图片 URL:{ctx}"))
52    }
53}
54
55/// 抓取 URL → 图片字节 → 走共享入库流水线 → 返回 `/uploads/...` 结果。
56///
57/// `original_filename` 从 URL 路径末段推导(无则 None)。见模块头部的 SSRF 防护说明。
58pub(crate) async fn fetch_and_ingest(url: &str) -> Result<UploadOutcome, FetchError> {
59    // 1. 解析 URL:强制 https。
60    let uri: Uri = url
61        .parse()
62        .map_err(|e| FetchError::invalid(e, "URL 格式"))?;
63    let scheme = uri.scheme_str().unwrap_or("");
64    if scheme != "https" {
65        return Err(FetchError::Invalid("仅支持 https:// 图片 URL".into()));
66    }
67    let host = uri
68        .host()
69        .ok_or_else(|| FetchError::Invalid("URL 缺少主机名".into()))?;
70    if host.is_empty() {
71        return Err(FetchError::Invalid("URL 缺少主机名".into()));
72    }
73
74    // 2. 解析主机名 → IP 列表,校验每一个都非私网/保留段。
75    //    ToSocketAddrs 是阻塞的(getaddrinfo),放进 spawn_blocking。
76    let host_owned = host.to_string();
77    let port = uri.port_u16().unwrap_or(443);
78    let resolve_target = format!("{host_owned}:{port}");
79    let addrs = tokio::task::spawn_blocking(move || {
80        resolve_target
81            .to_socket_addrs()
82            .map(|i| i.collect::<Vec<SocketAddr>>())
83    })
84    .await
85    .map_err(|e| FetchError::invalid(e, "解析任务失败"))?
86    .map_err(|e| FetchError::invalid(e, "DNS 解析失败"))?;
87
88    if addrs.is_empty() {
89        return Err(FetchError::Invalid("DNS 解析无结果".into()));
90    }
91
92    // 全部 IP 必须通过校验;任一私网 IP 即拒绝(防多记录里夹带内网)。
93    for addr in &addrs {
94        let ip = addr.ip();
95        if is_forbidden_ip(&ip) {
96            tracing::warn!("url_fetch SSRF 拒绝:{} → {}", host_owned, ip);
97            return Err(FetchError::Invalid("目标地址位于禁止的网络段".into()));
98        }
99    }
100    // 钉死首个通过校验的 IP,杜绝 reqwest 二次 DNS(rebinding)。
101    let locked_ip = addrs[0].ip();
102
103    // 3. 构建 reqwest 客户端:禁重定向 + 锁 IP + 超时 + UA。
104    let client = reqwest::Client::builder()
105        .redirect(reqwest::redirect::Policy::none())
106        .resolve(&host_owned, SocketAddr::from((locked_ip, port)))
107        .timeout(Duration::from_secs(30))
108        .connect_timeout(Duration::from_secs(10))
109        .user_agent("yggdrasil-mcp-media/1.0")
110        .build()
111        .map_err(|e| {
112            tracing::error!("url_fetch client build: {e}");
113            FetchError::Fetch("构建抓取客户端失败")
114        })?;
115
116    // 4. 发起 GET。
117    let resp = client.get(url).send().await.map_err(|e| {
118        tracing::warn!("url_fetch GET {url} failed: {e}");
119        FetchError::Fetch("抓取图片失败")
120    })?;
121
122    if !resp.status().is_success() {
123        return Err(FetchError::BadStatus(format!("远端返回 {}", resp.status())));
124    }
125
126    // 5. 流式读取,累计字节超 MAX_FILE_SIZE 立即中止。
127    //    Content-Length 可伪造/缺省,不能信任;以累计字节为准。
128    let mut total: usize = 0;
129    let mut buf: Vec<u8> = Vec::new();
130    let mut stream = resp.bytes_stream();
131    use futures::StreamExt;
132    while let Some(chunk) = stream.next().await {
133        let chunk = chunk.map_err(|e| {
134            tracing::warn!("url_fetch stream read failed: {e}");
135            FetchError::Fetch("读取图片数据失败")
136        })?;
137        total += chunk.len();
138        if total > MAX_FILE_SIZE {
139            return Err(FetchError::TooLarge);
140        }
141        buf.extend_from_slice(&chunk);
142    }
143
144    let data = Bytes::from(buf);
145
146    // 6. 从 URL 路径末段推导展示文件名(仅 assets 表展示字段,不影响落盘)。
147    let original_filename = uri
148        .path()
149        .rsplit('/')
150        .next()
151        .filter(|s| !s.is_empty())
152        .map(|s| s.to_string());
153
154    // 7. 走共享入库流水线(magic bytes 二次验真 + 尺寸 + 去重 + 转码 + 落盘)。
155    process_image_upload(data, original_filename)
156        .await
157        .map_err(|e| match e {
158            UploadError::TooLarge => FetchError::TooLarge,
159            other => {
160                let ctx = match other {
161                    UploadError::Empty => "空响应",
162                    UploadError::BadType => "非图片或格式不支持",
163                    UploadError::Oversized => "图片尺寸超限",
164                    UploadError::Corrupt => "图片损坏",
165                    UploadError::TooLarge => "文件过大",
166                    UploadError::Internal(_) => "入库失败",
167                };
168                FetchError::BadStatus(ctx.into())
169            }
170        })
171}
172
173/// SSRF 拒绝:私网 / 回环 / 链路本地 / 保留 / CGNAT 等非公网单播段。
174fn is_forbidden_ip(ip: &IpAddr) -> bool {
175    match ip {
176        IpAddr::V4(v4) => {
177            v4.is_private()        // 10/8, 172.16/12, 192.168/16
178            || v4.is_loopback()    // 127/8
179            || v4.is_link_local()  // 169.254/16
180            || v4.is_unspecified() // 0.0.0.0
181            || v4.is_broadcast()   // 255.255.255.255
182            || v4.is_documentation() // 192.0.2/24, 198.51.100/24, 203.0.113/24
183            // CGNAT 100.64.0.0/10(std 未覆盖,手动判)
184            || {
185                let o = v4.octets();
186                o[0] == 100 && o[1] >= 64 && o[1] <= 127
187            }
188        }
189        IpAddr::V6(v6) => {
190            v6.is_loopback()       // ::1
191            || v6.is_unspecified() // ::
192            || v6.is_multicast()   // ff00::/8
193            || {
194                let s = v6.segments();
195                (s[0] & 0xfe00) == 0xfc00 // ULA fc00::/7
196                || (s[0] & 0xffc0) == 0xfe80 // 链路本地 fe80::/10
197            }
198        }
199    }
200}
201
202#[cfg(all(test, feature = "server"))]
203mod tests {
204    use super::is_forbidden_ip;
205    use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
206
207    #[test]
208    fn forbidden_ipv4_ranges() {
209        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1))));
210        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1))));
211        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(172, 16, 0, 1))));
212        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1))));
213        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(169, 254, 1, 1))));
214        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0))));
215        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(100, 64, 0, 1)))); // CGNAT
216        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(
217            100, 127, 255, 255
218        )))); // CGNAT 末
219        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(192, 0, 2, 1)))); // 文档段
220    }
221
222    #[test]
223    fn allowed_ipv4_ranges() {
224        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
225        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1))));
226        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(99, 63, 0, 1)))); // CGNAT 前
227        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(
228            100, 63, 255, 255
229        )))); // CGNAT 前
230        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(100, 128, 0, 1)))); // CGNAT 后
231        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(
232            172, 15, 255, 255
233        )))); // 172.16 前
234        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(172, 32, 0, 1)))); // 172.16/12 后
235    }
236
237    #[test]
238    fn forbidden_ipv6_ranges() {
239        assert!(is_forbidden_ip(&IpAddr::V6(Ipv6Addr::LOCALHOST))); // ::1
240        assert!(is_forbidden_ip(&IpAddr::V6(Ipv6Addr::UNSPECIFIED))); // ::
241        assert!(is_forbidden_ip(&IpAddr::V6("ff02::1".parse().unwrap()))); // 多播
242        assert!(is_forbidden_ip(&IpAddr::V6("fc00::1".parse().unwrap()))); // ULA
243        assert!(is_forbidden_ip(&IpAddr::V6("fd12::1".parse().unwrap()))); // ULA
244        assert!(is_forbidden_ip(&IpAddr::V6("fe80::1".parse().unwrap()))); // 链路本地
245    }
246
247    #[test]
248    fn allowed_ipv6_ranges() {
249        assert!(!is_forbidden_ip(&IpAddr::V6(
250            "2606:4700::1".parse().unwrap()
251        ))); // Cloudflare 公网
252        assert!(!is_forbidden_ip(&IpAddr::V6(
253            "2001:4860:4860::8888".parse().unwrap()
254        ))); // Google DNS
255    }
256}