XFE Git
XFE Studio Git
Git 首页 全局搜索
XFE 主站 文档 NuGet
公开
关注 0 Fork 0 Star 0
返回提交历史

XFEstudio/gpt4free

feat: Enhance URL safety checks and integrate is_safe_url across image handling #3397

b6a6db2c
hlohaus <983577+hlohaus@users.noreply.github.com>
提交于

代码差异

4 个文件 +73 -30
Modified etc/unittest/backend.py +20 -1
@@ -2,7 +2,8 @@ from __future__ import annotations
2 2
3 3 import unittest
4 4 import asyncio
5 from unittest.mock import MagicMock
5 import socket
6 from unittest.mock import MagicMock, patch
6 7 from g4f.errors import MissingRequirementsError
7 8 try:
8 9 from g4f.gui.server.backend_api import Backend_Api
@@ -43,6 +44,24 @@ class TestBackendApi(unittest.TestCase):
43 44 self.assertIsInstance(response, list)
44 45 self.assertTrue(len(response) > 0)
45 46
47 @patch('g4f.gui.server.backend_api.socket.getaddrinfo')
48 def test_is_safe_url_with_backslash_confusion(self, mock_getaddrinfo):
49 mock_getaddrinfo.return_value = [(socket.AF_INET, socket.SOCK_STREAM, 6, '', ('127.0.0.1', 0))]
50 from g4f.gui.server.backend_api import _is_safe_url
51 self.assertFalse(_is_safe_url('http://127.0.0.1:6666\\@www.baidu.com'))
52
53 @patch('g4f.gui.server.backend_api.socket.getaddrinfo')
54 def test_is_safe_url_blocks_private(self, mock_getaddrinfo):
55 mock_getaddrinfo.return_value = [(socket.AF_INET, socket.SOCK_STREAM, 6, '', ('127.0.0.1', 0))]
56 from g4f.gui.server.backend_api import _is_safe_url
57 self.assertFalse(_is_safe_url('http://127.0.0.1'))
58
59 @patch('g4f.gui.server.backend_api.socket.getaddrinfo')
60 def test_is_safe_url_allows_public(self, mock_getaddrinfo):
61 mock_getaddrinfo.return_value = [(socket.AF_INET, socket.SOCK_STREAM, 6, '', ('8.8.8.8', 0))]
62 from g4f.gui.server.backend_api import _is_safe_url
63 self.assertTrue(_is_safe_url('http://example.com'))
64
46 65 def test_search(self):
47 66 if not has_search:
48 67 self.skipTest("import error")
Modified g4f/gui/server/backend_api.py +2 -28
@@ -11,10 +11,8 @@ import asyncio
11 11 import shutil
12 12 import random
13 13 import datetime
14 import ipaddress
15 import socket
16 14 from hashlib import sha256
17 from urllib.parse import quote_plus, urlparse
15 from urllib.parse import quote_plus
18 16 from functools import lru_cache
19 17 from flask import Flask, Response, redirect, request, jsonify, send_from_directory
20 18 from werkzeug.exceptions import NotFound
@@ -46,7 +44,7 @@ from ...client.helper import filter_markdown
46 44 from ...tools.files import supports_filename, get_streaming, get_bucket_dir, get_tempfile
47 45 from ...tools.run_tools import iter_run_tools
48 46 from ...errors import ModelNotFoundError, ProviderNotFoundError, MissingAuthError, RateLimitError
49 from ...image import is_allowed_extension, process_image, MEDIA_TYPE_MAP
47 from ...image import is_allowed_extension, process_image, MEDIA_TYPE_MAP, is_safe_url as _is_safe_url
50 48 from ...cookies import get_cookies_dir
51 49 from ...image.copy_images import secure_filename, get_source_url, get_media_dir, copy_media
52 50 from ...client.service import get_model_and_provider
@@ -60,30 +58,6 @@ logger = logging.getLogger(__name__)
60 58
61 59 _DATE_RE = re.compile(r'^\d{4}-\d{2}-\d{2}$')
62 60
63 def _is_safe_url(url: str) -> bool:
64 """Return True only for http/https URLs that do not point to private/loopback/reserved addresses."""
65 try:
66 parsed = urlparse(url)
67 if parsed.scheme not in ("http", "https"):
68 return False
69 hostname = parsed.hostname
70 if hostname is None:
71 return False
72 # Resolve all IP addresses for the hostname and reject if any is non-public.
73 # Validating all addresses reduces the window for DNS rebinding attacks.
74 addr_infos = socket.getaddrinfo(hostname, None)
75 if not addr_infos:
76 return False
77 for addr_info in addr_infos:
78 addr = ipaddress.ip_address(addr_info[4][0])
79 if (addr.is_private or addr.is_loopback or addr.is_link_local
80 or addr.is_reserved or addr.is_multicast or addr.is_unspecified):
81 return False
82 except Exception as e:
83 logger.debug("URL safety check failed for %r: %s", url, e)
84 return False
85 return True
86
87 61 def safe_iter_generator(generator: Generator) -> Generator:
88 62 start = next(generator)
89 63 def iter_generator():
Modified g4f/image/__init__.py +48 -0
@@ -4,6 +4,8 @@ import os
4 4 import re
5 5 import io
6 6 import base64
7 import socket
8 import ipaddress
7 9 from io import BytesIO
8 10 from pathlib import Path
9 11 from typing import Optional
@@ -11,6 +13,11 @@ from urllib.parse import urlparse
11 13
12 14 import requests
13 15
16 try:
17 from urllib3.util import parse_url as urllib3_parse_url
18 except ImportError:
19 urllib3_parse_url = None
20
14 21 try:
15 22 from PIL import Image, ImageOps
16 23 has_requirements = True
@@ -103,6 +110,45 @@ def is_allowed_extension(filename: str) -> Optional[str]:
103 110 return None
104 111 return EXTENSIONS_MAP[extension]
105 112
113
114 def is_safe_url(url: str) -> bool:
115 """Return True only for http/https URLs that do not point to private/loopback/reserved addresses."""
116 try:
117 parsed = urlparse(url)
118
119 if parsed.scheme not in ("http", "https"):
120 return False
121
122 if "\\" in url:
123 return False
124
125 hostname = parsed.hostname
126 if hostname is None:
127 return False
128
129 if urllib3_parse_url is not None:
130 parsed_urllib3 = urllib3_parse_url(url)
131 if parsed_urllib3.host and parsed_urllib3.host != hostname:
132 return False
133 hostname = parsed_urllib3.host or hostname
134
135 if hostname is None:
136 return False
137
138 addr_infos = socket.getaddrinfo(hostname, None)
139 if not addr_infos:
140 return False
141
142 for addr_info in addr_infos:
143 addr = ipaddress.ip_address(addr_info[4][0])
144 if (addr.is_private or addr.is_loopback or addr.is_link_local
145 or addr.is_reserved or addr.is_multicast or addr.is_unspecified):
146 return False
147 except Exception:
148 return False
149 return True
150
151
106 152 def is_data_an_media(data, filename: str = None) -> str:
107 153 content_type = is_data_an_audio(data, filename)
108 154 if content_type is not None:
@@ -378,6 +424,8 @@ def to_bytes(image: ImageType) -> bytes:
378 424 is_data_uri_an_image(image)
379 425 return extract_data_uri(image)
380 426 elif image.startswith("http://") or image.startswith("https://"):
427 if not is_safe_url(image):
428 raise ValueError("Invalid or disallowed media URL")
381 429 path: str = urlparse(image).path
382 430 if path.startswith("/files/"):
383 431 path = get_bucket_dir(*path.split("/")[2:])
Modified g4f/image/copy_images.py +3 -1
@@ -16,7 +16,7 @@ from ..requests.aiohttp import get_connector
16 16 from ..image import MEDIA_TYPE_MAP, EXTENSIONS_MAP
17 17 from ..tools.files import secure_filename
18 18 from ..providers.response import ImageResponse, AudioResponse, VideoResponse, quote_url
19 from . import is_accepted_format, extract_data_uri
19 from . import is_accepted_format, extract_data_uri, is_safe_url
20 20 from .. import debug
21 21
22 22 # Directory for storing generated media files
@@ -170,6 +170,8 @@ async def copy_media(
170 170 with open(target_path, "wb") as f:
171 171 f.write(extract_data_uri(image))
172 172 elif not os.path.exists(target_path) or os.lstat(target_path).st_size <= 0:
173 if not is_safe_url(image):
174 raise ValueError("Invalid or disallowed media URL")
173 175 # Use aiohttp to fetch the image
174 176 async with session.get(image, ssl=ssl) as response:
175 177 response.raise_for_status()