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(
59    url: &str,
60    alt: Option<String>,
61) -> Result<UploadOutcome, FetchError> {
62    // 1. 解析 URL:强制 https。
63    let uri: Uri = url
64        .parse()
65        .map_err(|e| FetchError::invalid(e, "URL 格式"))?;
66    let scheme = uri.scheme_str().unwrap_or("");
67    if scheme != "https" {
68        return Err(FetchError::Invalid("仅支持 https:// 图片 URL".into()));
69    }
70    let host = uri
71        .host()
72        .ok_or_else(|| FetchError::Invalid("URL 缺少主机名".into()))?;
73    if host.is_empty() {
74        return Err(FetchError::Invalid("URL 缺少主机名".into()));
75    }
76
77    // 2. 解析主机名 → IP 列表,校验每一个都非私网/保留段。
78    //    ToSocketAddrs 是阻塞的(getaddrinfo),放进 spawn_blocking。
79    let host_owned = host.to_string();
80    let port = uri.port_u16().unwrap_or(443);
81    let resolve_target = format!("{host_owned}:{port}");
82    let addrs = tokio::task::spawn_blocking(move || {
83        resolve_target
84            .to_socket_addrs()
85            .map(|i| i.collect::<Vec<SocketAddr>>())
86    })
87    .await
88    .map_err(|e| FetchError::invalid(e, "解析任务失败"))?
89    .map_err(|e| FetchError::invalid(e, "DNS 解析失败"))?;
90
91    if addrs.is_empty() {
92        return Err(FetchError::Invalid("DNS 解析无结果".into()));
93    }
94
95    // 全部 IP 必须通过校验;任一私网 IP 即拒绝(防多记录里夹带内网)。
96    for addr in &addrs {
97        let ip = addr.ip();
98        if is_forbidden_ip(&ip) {
99            tracing::warn!("url_fetch SSRF 拒绝:{} → {}", host_owned, ip);
100            return Err(FetchError::Invalid("目标地址位于禁止的网络段".into()));
101        }
102    }
103    // 钉死首个通过校验的 IP,杜绝 reqwest 二次 DNS(rebinding)。
104    let locked_ip = addrs[0].ip();
105
106    // 3. 构建 reqwest 客户端:禁重定向 + 锁 IP + 超时 + UA。
107    let client = reqwest::Client::builder()
108        .redirect(reqwest::redirect::Policy::none())
109        .resolve(&host_owned, SocketAddr::from((locked_ip, port)))
110        .timeout(Duration::from_secs(30))
111        .connect_timeout(Duration::from_secs(10))
112        .user_agent("yggdrasil-mcp-media/1.0")
113        .build()
114        .map_err(|e| {
115            tracing::error!("url_fetch client build: {e}");
116            FetchError::Fetch("构建抓取客户端失败")
117        })?;
118
119    // 4. 发起 GET。
120    let resp = client.get(url).send().await.map_err(|e| {
121        tracing::warn!("url_fetch GET {url} failed: {e}");
122        FetchError::Fetch("抓取图片失败")
123    })?;
124
125    if !resp.status().is_success() {
126        return Err(FetchError::BadStatus(format!("远端返回 {}", resp.status())));
127    }
128
129    // 5. 流式读取,累计字节超 MAX_FILE_SIZE 立即中止。
130    //    Content-Length 可伪造/缺省,不能信任;以累计字节为准。
131    let mut total: usize = 0;
132    let mut buf: Vec<u8> = Vec::new();
133    let mut stream = resp.bytes_stream();
134    use futures::StreamExt;
135    while let Some(chunk) = stream.next().await {
136        let chunk = chunk.map_err(|e| {
137            tracing::warn!("url_fetch stream read failed: {e}");
138            FetchError::Fetch("读取图片数据失败")
139        })?;
140        total += chunk.len();
141        if total > MAX_FILE_SIZE {
142            return Err(FetchError::TooLarge);
143        }
144        buf.extend_from_slice(&chunk);
145    }
146
147    let data = Bytes::from(buf);
148
149    // 6. 从 URL 路径末段推导展示文件名(仅 assets 表展示字段,不影响落盘)。
150    let original_filename = uri
151        .path()
152        .rsplit('/')
153        .next()
154        .filter(|s| !s.is_empty())
155        .map(|s| s.to_string());
156
157    // 7. 走共享入库流水线(magic bytes 二次验真 + 尺寸 + 去重 + 转码 + 落盘)。
158    process_image_upload(data, original_filename, alt)
159        .await
160        .map_err(|e| match e {
161            UploadError::TooLarge => FetchError::TooLarge,
162            other => {
163                let ctx = match other {
164                    UploadError::Empty => "空响应",
165                    UploadError::BadType => "非图片或格式不支持",
166                    UploadError::Oversized => "图片尺寸超限",
167                    UploadError::Corrupt => "图片损坏",
168                    UploadError::TooLarge => "文件过大",
169                    UploadError::Internal(_) => "入库失败",
170                };
171                FetchError::BadStatus(ctx.into())
172            }
173        })
174}
175
176/// SSRF 拒绝:私网 / 回环 / 链路本地 / 保留 / CGNAT 等非公网单播段。
177fn is_forbidden_ip(ip: &IpAddr) -> bool {
178    match ip {
179        IpAddr::V4(v4) => {
180            v4.is_private()        // 10/8, 172.16/12, 192.168/16
181            || v4.is_loopback()    // 127/8
182            || v4.is_link_local()  // 169.254/16
183            || v4.is_unspecified() // 0.0.0.0
184            || v4.is_broadcast()   // 255.255.255.255
185            || v4.is_documentation() // 192.0.2/24, 198.51.100/24, 203.0.113/24
186            // CGNAT 100.64.0.0/10(std 未覆盖,手动判)
187            || {
188                let o = v4.octets();
189                o[0] == 100 && o[1] >= 64 && o[1] <= 127
190            }
191        }
192        IpAddr::V6(v6) => {
193            if v6.is_loopback() || v6.is_unspecified() {
194                return true;
195            }
196            // IPv4-mapped / compatible IPv6 地址仍可能连接到 IPv4 内网服务;
197            // 先归一化后复用完整 IPv4 禁止网段检查,避免 ::ffff:127.0.0.1 绕过。
198            if let Some(v4) = v6.to_ipv4() {
199                return is_forbidden_ip(&IpAddr::V4(v4));
200            }
201            v6.is_multicast()   // ff00::/8
202            || {
203                let s = v6.segments();
204                (s[0] & 0xfe00) == 0xfc00 // ULA fc00::/7
205                || (s[0] & 0xffc0) == 0xfe80 // 链路本地 fe80::/10
206            }
207        }
208    }
209}
210
211#[cfg(all(test, feature = "server"))]
212mod tests {
213    use super::is_forbidden_ip;
214    use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
215
216    #[test]
217    fn forbidden_ipv4_ranges() {
218        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1))));
219        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1))));
220        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(172, 16, 0, 1))));
221        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1))));
222        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(169, 254, 1, 1))));
223        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0))));
224        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(100, 64, 0, 1)))); // CGNAT
225        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(
226            100, 127, 255, 255
227        )))); // CGNAT 末
228        assert!(is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(192, 0, 2, 1)))); // 文档段
229    }
230
231    #[test]
232    fn allowed_ipv4_ranges() {
233        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
234        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1))));
235        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(99, 63, 0, 1)))); // CGNAT 前
236        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(
237            100, 63, 255, 255
238        )))); // CGNAT 前
239        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(100, 128, 0, 1)))); // CGNAT 后
240        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(
241            172, 15, 255, 255
242        )))); // 172.16 前
243        assert!(!is_forbidden_ip(&IpAddr::V4(Ipv4Addr::new(172, 32, 0, 1)))); // 172.16/12 后
244    }
245
246    #[test]
247    fn forbidden_ipv6_ranges() {
248        assert!(is_forbidden_ip(&IpAddr::V6(Ipv6Addr::LOCALHOST))); // ::1
249        assert!(is_forbidden_ip(&IpAddr::V6(Ipv6Addr::UNSPECIFIED))); // ::
250        assert!(is_forbidden_ip(&IpAddr::V6("ff02::1".parse().unwrap()))); // 多播
251        assert!(is_forbidden_ip(&IpAddr::V6("fc00::1".parse().unwrap()))); // ULA
252        assert!(is_forbidden_ip(&IpAddr::V6("fd12::1".parse().unwrap()))); // ULA
253        assert!(is_forbidden_ip(&IpAddr::V6("fe80::1".parse().unwrap()))); // 链路本地
254    }
255
256    #[test]
257    fn forbidden_ipv4_mapped_ipv6_ranges() {
258        assert!(is_forbidden_ip(&IpAddr::V6(
259            "::ffff:127.0.0.1".parse().expect("valid mapped loopback")
260        )));
261        assert!(is_forbidden_ip(&IpAddr::V6(
262            "::ffff:169.254.169.254"
263                .parse()
264                .expect("valid mapped link-local")
265        )));
266        assert!(is_forbidden_ip(&IpAddr::V6(
267            "::ffff:10.0.0.1".parse().expect("valid mapped private")
268        )));
269    }
270
271    #[test]
272    fn allowed_ipv6_ranges() {
273        assert!(!is_forbidden_ip(&IpAddr::V6(
274            "2606:4700::1".parse().unwrap()
275        ))); // Cloudflare 公网
276        assert!(!is_forbidden_ip(&IpAddr::V6(
277            "2001:4860:4860::8888".parse().unwrap()
278        ))); // Google DNS
279    }
280}