| 1 | """S3 URL parsing helpers for server-side object downloads.""" |
| 2 | |
| 3 | from __future__ import annotations |
| 4 | |
| 5 | |
| 6 | class S3UrlError(ValueError): |
| 7 | pass |
| 8 | |
| 9 | |
| 10 | def is_s3_url(url: str) -> bool: |
| 11 | return url.strip().lower().startswith("s3://") |
| 12 | |
| 13 | |
| 14 | def parse_s3_bucket_and_key(s3_url: str) -> tuple[str, str]: |
| 15 | """Parse ``s3://bucket/key/path`` into ``(bucket, key)``.""" |
| 16 | raw = s3_url.strip() |
| 17 | if not is_s3_url(raw): |
| 18 | raise S3UrlError(f"不是有效的S3 URL: {s3_url}") |
| 19 | without_scheme = raw[5:] |
| 20 | parts = without_scheme.split("/", 1) |
| 21 | if len(parts) < 2 or not parts[0].strip() or not parts[1].strip(): |
| 22 | raise S3UrlError(f"S3 URL格式错误: {s3_url}") |
| 23 | return parts[0].strip(), parts[1].strip() |
| 24 | |
| 25 | |
| 26 | def validate_s3_url(s3_url: str, *, allowed_bucket: str) -> str: |
| 27 | """Return object key when *s3_url* targets *allowed_bucket*.""" |
| 28 | bucket, key = parse_s3_bucket_and_key(s3_url) |
| 29 | if bucket != allowed_bucket: |
| 30 | raise S3UrlError(f"不允许访问该存储桶: {bucket}") |
| 31 | if ".." in key.split("/"): |
| 32 | raise S3UrlError("非法的S3路径") |
| 33 | return key |
| 34 |