|
38 | 38 | from typing_extensions import Self |
39 | 39 |
|
40 | 40 | _TEMP_DIR_PREFIX = "flink-agents-skills-" |
| 41 | +MAX_DOWNLOAD_BYTES: int = 512 * 1024 * 1024 |
| 42 | +MAX_EXTRACT_ENTRY_BYTES: int = 200 * 1024 * 1024 |
| 43 | +MAX_EXTRACT_TOTAL_BYTES: int = 1024 * 1024 * 1024 |
| 44 | +MAX_EXTRACT_ENTRIES: int = 10_000 |
41 | 45 | logger = logging.getLogger(__name__) |
42 | 46 |
|
43 | 47 |
|
@@ -179,15 +183,101 @@ def extract_zip_safely(zip_path: Path) -> Materialized: |
179 | 183 | # Construct the handle before validation so the (empty) tempdir is always reclaimed, |
180 | 184 | # even if validation raises. |
181 | 185 | materialized = Materialized(extract_dir) |
182 | | - with zipfile.ZipFile(zip_path) as zf: |
183 | | - for member in zf.infolist(): |
184 | | - target = (extract_dir / member.filename).resolve() |
185 | | - if not target.is_relative_to(extract_dir): |
186 | | - msg = f"Unsafe zip entry: {member.filename}" |
187 | | - raise ValueError(msg) |
188 | | - zf.extractall(extract_dir) |
| 186 | + try: |
| 187 | + _extract_zip_to_dir(zip_path, extract_dir) |
| 188 | + except Exception: |
| 189 | + materialized.close() |
| 190 | + raise |
189 | 191 | return materialized |
190 | 192 |
|
| 193 | +def _validate_zip_members(members: list, extract_dir: Path) -> None: |
| 194 | + if len(members) > MAX_EXTRACT_ENTRIES: |
| 195 | + msg = ( |
| 196 | + f"Skill archive contains {len(members)} entries, " |
| 197 | + f"exceeding the limit of {MAX_EXTRACT_ENTRIES}" |
| 198 | + ) |
| 199 | + raise ValueError(msg) |
| 200 | + |
| 201 | + for member in members: |
| 202 | + target = (extract_dir / member.filename).resolve() |
| 203 | + if not target.is_relative_to(extract_dir): |
| 204 | + msg = f"Unsafe zip entry: {member.filename}" |
| 205 | + raise ValueError(msg) |
| 206 | + |
| 207 | + total_declared = 0 |
| 208 | + for member in members: |
| 209 | + if member.is_dir(): |
| 210 | + continue |
| 211 | + declared = member.file_size |
| 212 | + if declared > MAX_EXTRACT_ENTRY_BYTES: |
| 213 | + msg = ( |
| 214 | + f"Skill archive entry '{member.filename}' declared size {declared} " |
| 215 | + f"exceeds the per-entry limit of {MAX_EXTRACT_ENTRY_BYTES} bytes" |
| 216 | + ) |
| 217 | + raise ValueError(msg) |
| 218 | + if declared > 0: |
| 219 | + total_declared += declared |
| 220 | + if total_declared > MAX_EXTRACT_TOTAL_BYTES: |
| 221 | + msg = ( |
| 222 | + f"Skill archive declared total uncompressed size {total_declared} " |
| 223 | + f"exceeds the limit of {MAX_EXTRACT_TOTAL_BYTES} bytes" |
| 224 | + ) |
| 225 | + raise ValueError(msg) |
| 226 | + |
| 227 | +def _extract_zip_to_dir(zip_path: Path, extract_dir: Path) -> None: |
| 228 | + with zipfile.ZipFile(zip_path) as zf: |
| 229 | + members = zf.infolist() |
| 230 | + _validate_zip_members(members, extract_dir) |
| 231 | + buf = bytearray(65536) |
| 232 | + total_written = 0 |
| 233 | + for member in members: |
| 234 | + target = (extract_dir / member.filename).resolve() |
| 235 | + if member.is_dir(): |
| 236 | + target.mkdir(parents=True, exist_ok=True) |
| 237 | + continue |
| 238 | + target.parent.mkdir(parents=True, exist_ok=True) |
| 239 | + per_entry_written = 0 |
| 240 | + with zf.open(member) as src, target.open("xb") as dst: |
| 241 | + while True: |
| 242 | + n = src.readinto(buf) |
| 243 | + if not n: |
| 244 | + break |
| 245 | + _check_entry_size(member.filename, per_entry_written, n) |
| 246 | + _check_total_size(total_written, n) |
| 247 | + dst.write(buf[:n]) |
| 248 | + per_entry_written += n |
| 249 | + total_written += n |
| 250 | + |
| 251 | +def _check_entry_size(filename: str, already_written: int, chunk: int) -> None: |
| 252 | + if already_written + chunk > MAX_EXTRACT_ENTRY_BYTES: |
| 253 | + msg = ( |
| 254 | + f"Skill archive entry '{filename}' exceeds the " |
| 255 | + f"per-entry limit of {MAX_EXTRACT_ENTRY_BYTES} bytes" |
| 256 | + ) |
| 257 | + raise ValueError(msg) |
| 258 | + |
| 259 | +def _check_total_size(already_written: int, chunk: int) -> None: |
| 260 | + if already_written + chunk > MAX_EXTRACT_TOTAL_BYTES: |
| 261 | + msg = ( |
| 262 | + f"Skill archive total extracted size exceeds the limit of " |
| 263 | + f"{MAX_EXTRACT_TOTAL_BYTES} bytes" |
| 264 | + ) |
| 265 | + raise ValueError(msg) |
| 266 | + |
| 267 | +def _check_declared_download_size(content_length: int | None) -> None: |
| 268 | + if content_length is not None and content_length > MAX_DOWNLOAD_BYTES: |
| 269 | + msg = ( |
| 270 | + f"Skill archive download size declared as {content_length} bytes, " |
| 271 | + f"exceeding the limit of {MAX_DOWNLOAD_BYTES} bytes" |
| 272 | + ) |
| 273 | + raise ValueError(msg) |
| 274 | + |
| 275 | +def _check_download_size(already_written: int, chunk: int) -> None: |
| 276 | + if already_written + chunk > MAX_DOWNLOAD_BYTES: |
| 277 | + msg = f"Skill archive download exceeded the limit of {MAX_DOWNLOAD_BYTES} bytes" |
| 278 | + raise ValueError(msg) |
| 279 | + |
| 280 | + |
191 | 281 |
|
192 | 282 | def download_to_tempfile( |
193 | 283 | url: str, timeout: int = 90, *, allow_insecure_http: bool = False |
@@ -239,7 +329,21 @@ def download_to_tempfile( |
239 | 329 | redact_skill_url(url), |
240 | 330 | redact_skill_url(final_url), |
241 | 331 | ) |
242 | | - shutil.copyfileobj(resp, out) |
| 332 | + raw_cl = resp.headers.get("Content-Length") |
| 333 | + if raw_cl is not None: |
| 334 | + try: |
| 335 | + content_length = int(raw_cl) |
| 336 | + except ValueError: |
| 337 | + content_length = None |
| 338 | + _check_declared_download_size(content_length) |
| 339 | + written = 0 |
| 340 | + while True: |
| 341 | + chunk = resp.read(65536) |
| 342 | + if not chunk: |
| 343 | + break |
| 344 | + _check_download_size(written, len(chunk)) |
| 345 | + out.write(chunk) |
| 346 | + written += len(chunk) |
243 | 347 | except Exception: |
244 | 348 | tmp_path.unlink(missing_ok=True) |
245 | 349 | raise |
|
0 commit comments