返回提交历史
Modified
etc/unittest/backend.py
+20
-1
Modified
g4f/gui/server/backend_api.py
+2
-28
Modified
g4f/image/__init__.py
+48
-0
Modified
g4f/image/copy_images.py
+3
-1
XFEstudio/gpt4free
feat: Enhance URL safety checks and integrate is_safe_url across image handling #3397
b6a6db2c
代码差异
4 个文件
+73
-30
@@ -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")
@@ -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():
@@ -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:])
@@ -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()