tokens and more
This commit is contained in:
Executable
+829
@@ -0,0 +1,829 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Populate RustPad through its HTTP API with generated test data."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import concurrent.futures
|
||||
import getpass
|
||||
from html.parser import HTMLParser
|
||||
import json
|
||||
import mimetypes
|
||||
import os
|
||||
import random
|
||||
import ssl
|
||||
import string
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from http.cookies import SimpleCookie
|
||||
from typing import Any, Iterable
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.parse import quote, urljoin, urlparse
|
||||
from urllib.request import HTTPSHandler, Request, build_opener
|
||||
|
||||
UNSAFE_METHODS = {"POST", "PUT", "PATCH", "DELETE"}
|
||||
WIKIPEDIA_RANDOM_URL = "https://en.wikipedia.org/wiki/Special:Random"
|
||||
MAX_SOURCE_BYTES = 1_500_000
|
||||
MAX_DOCUMENT_BYTES = 1_800_000
|
||||
MAX_WIKIPEDIA_ATTACHMENT_BYTES = 5_000_000
|
||||
WIKIMEDIA_IMAGE_HOST_SUFFIX = ".wikimedia.org"
|
||||
WIKIPEDIA_IMAGE_EXTENSIONS = (".avif", ".bmp", ".gif", ".jpeg", ".jpg", ".png", ".webp")
|
||||
WIKIPEDIA_IMAGE_MIME_TYPES = {
|
||||
"image/avif",
|
||||
"image/bmp",
|
||||
"image/gif",
|
||||
"image/jpeg",
|
||||
"image/png",
|
||||
"image/webp",
|
||||
}
|
||||
|
||||
|
||||
class ApiFailure(RuntimeError):
|
||||
def __init__(self, status: int, message: str, path: str) -> None:
|
||||
super().__init__(f"{status} {path}: {message}")
|
||||
self.status = status
|
||||
self.message = message
|
||||
self.path = path
|
||||
|
||||
|
||||
class RustPadClient:
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
*,
|
||||
timeout: float,
|
||||
retries: int,
|
||||
insecure: bool,
|
||||
) -> None:
|
||||
self.base_url = base_url.rstrip("/") + "/"
|
||||
self.timeout = timeout
|
||||
self.retries = retries
|
||||
self.cookies: dict[str, str] = {}
|
||||
self.csrf_token: str | None = None
|
||||
self.lock = threading.RLock()
|
||||
self.context = ssl._create_unverified_context() if insecure else ssl.create_default_context()
|
||||
self.local = threading.local()
|
||||
|
||||
def _opener(self) -> Any:
|
||||
opener = getattr(self.local, "opener", None)
|
||||
if opener is None:
|
||||
opener = build_opener(HTTPSHandler(context=self.context))
|
||||
self.local.opener = opener
|
||||
return opener
|
||||
|
||||
def _cookie_header(self) -> str:
|
||||
with self.lock:
|
||||
return "; ".join(f"{name}={value}" for name, value in self.cookies.items())
|
||||
|
||||
def _store_cookies(self, headers: Any) -> None:
|
||||
values = headers.get_all("Set-Cookie") or []
|
||||
if not values:
|
||||
return
|
||||
with self.lock:
|
||||
for raw in values:
|
||||
parsed = SimpleCookie()
|
||||
parsed.load(raw)
|
||||
for name, morsel in parsed.items():
|
||||
if morsel["max-age"] == "0" or not morsel.value:
|
||||
self.cookies.pop(name, None)
|
||||
else:
|
||||
self.cookies[name] = morsel.value
|
||||
|
||||
def refresh_csrf(self) -> str:
|
||||
with self.lock:
|
||||
data = self.request("GET", "/api/security/csrf", retry_csrf=False)
|
||||
token = str(data.get("token") or "")
|
||||
if not token:
|
||||
raise RuntimeError("The server did not return a CSRF token.")
|
||||
self.csrf_token = token
|
||||
return token
|
||||
|
||||
def _request_bytes(
|
||||
self,
|
||||
method: str,
|
||||
path: str,
|
||||
body: bytes | None,
|
||||
content_type: str | None,
|
||||
*,
|
||||
retry_csrf: bool = True,
|
||||
) -> dict[str, Any]:
|
||||
method = method.upper()
|
||||
url = urljoin(self.base_url, path.lstrip("/"))
|
||||
|
||||
for attempt in range(self.retries + 1):
|
||||
headers = {
|
||||
"Accept": "application/json",
|
||||
"User-Agent": "RustPad-random-data/1.0",
|
||||
}
|
||||
if content_type:
|
||||
headers["Content-Type"] = content_type
|
||||
cookie = self._cookie_header()
|
||||
if cookie:
|
||||
headers["Cookie"] = cookie
|
||||
if method in UNSAFE_METHODS:
|
||||
if not self.csrf_token:
|
||||
self.refresh_csrf()
|
||||
headers["X-Rustpad-CSRF"] = self.csrf_token or ""
|
||||
|
||||
request = Request(url, data=body, headers=headers, method=method)
|
||||
try:
|
||||
with self._opener().open(request, timeout=self.timeout) as response:
|
||||
self._store_cookies(response.headers)
|
||||
raw = response.read()
|
||||
if not raw:
|
||||
return {}
|
||||
return json.loads(raw.decode("utf-8"))
|
||||
except HTTPError as error:
|
||||
self._store_cookies(error.headers)
|
||||
raw = error.read()
|
||||
try:
|
||||
data = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
except (UnicodeDecodeError, json.JSONDecodeError):
|
||||
data = {}
|
||||
message = str(data.get("error") or error.reason or "Request failed")
|
||||
if (
|
||||
error.code == 403
|
||||
and retry_csrf
|
||||
and "security token" in message.lower()
|
||||
):
|
||||
with self.lock:
|
||||
self.csrf_token = None
|
||||
self.refresh_csrf()
|
||||
return self._request_bytes(
|
||||
method,
|
||||
path,
|
||||
body,
|
||||
content_type,
|
||||
retry_csrf=False,
|
||||
)
|
||||
if error.code in {429, 500, 502, 503, 504} and attempt < self.retries:
|
||||
retry_after = error.headers.get("Retry-After")
|
||||
try:
|
||||
delay = float(retry_after) if retry_after else min(10.0, 0.5 * (2**attempt))
|
||||
except ValueError:
|
||||
delay = min(10.0, 0.5 * (2**attempt))
|
||||
time.sleep(delay + random.random() * 0.25)
|
||||
continue
|
||||
raise ApiFailure(error.code, message, path) from error
|
||||
except (URLError, TimeoutError, json.JSONDecodeError) as error:
|
||||
if attempt < self.retries:
|
||||
time.sleep(min(10.0, 0.5 * (2**attempt)) + random.random() * 0.25)
|
||||
continue
|
||||
raise RuntimeError(f"Request to {path} failed: {error}") from error
|
||||
|
||||
raise RuntimeError(f"Request to {path} failed after retries.")
|
||||
|
||||
def request(
|
||||
self,
|
||||
method: str,
|
||||
path: str,
|
||||
payload: dict[str, Any] | None = None,
|
||||
*,
|
||||
retry_csrf: bool = True,
|
||||
) -> dict[str, Any]:
|
||||
body = None if payload is None else json.dumps(payload, ensure_ascii=False).encode("utf-8")
|
||||
content_type = "application/json" if body is not None else None
|
||||
return self._request_bytes(
|
||||
method,
|
||||
path,
|
||||
body,
|
||||
content_type,
|
||||
retry_csrf=retry_csrf,
|
||||
)
|
||||
|
||||
def upload_multipart(
|
||||
self,
|
||||
path: str,
|
||||
*,
|
||||
filename: str,
|
||||
data: bytes,
|
||||
mime_type: str,
|
||||
) -> dict[str, Any]:
|
||||
boundary = f"rustpad-{''.join(random.choices(string.ascii_letters + string.digits, k=32))}"
|
||||
chunks = [
|
||||
f"--{boundary}\r\n".encode("ascii"),
|
||||
b'Content-Disposition: form-data; name="access_token"\r\n\r\n',
|
||||
b"\r\n",
|
||||
f"--{boundary}\r\n".encode("ascii"),
|
||||
(
|
||||
"Content-Disposition: form-data; name=\"file\"; "
|
||||
f"filename=\"{filename.replace(chr(34), '_')}\"\r\n"
|
||||
).encode("utf-8"),
|
||||
f"Content-Type: {mime_type}\r\n\r\n".encode("ascii"),
|
||||
data,
|
||||
b"\r\n",
|
||||
f"--{boundary}--\r\n".encode("ascii"),
|
||||
]
|
||||
return self._request_bytes(
|
||||
"POST",
|
||||
path,
|
||||
b"".join(chunks),
|
||||
f"multipart/form-data; boundary={boundary}",
|
||||
)
|
||||
|
||||
def login(self, user: str, password: str) -> dict[str, Any]:
|
||||
self.refresh_csrf()
|
||||
session = self.request(
|
||||
"POST",
|
||||
"/api/auth/login",
|
||||
{"email": user, "password": password},
|
||||
)
|
||||
verified = self.request("GET", "/api/auth/me")
|
||||
return verified or session
|
||||
|
||||
def create_pad(self, name: str, content: str) -> dict[str, Any]:
|
||||
return self.request("POST", "/api/pads", {"name": name, "content": content})
|
||||
|
||||
def create_workspace(self, name: str) -> dict[str, Any]:
|
||||
return self.request("POST", "/api/workspaces", {"name": name})
|
||||
|
||||
def create_workspace_note(self, workspace_slug: str, name: str, content: str) -> dict[str, Any]:
|
||||
return self.request(
|
||||
"POST",
|
||||
f"/api/workspaces/{quote(workspace_slug, safe='')}/notes",
|
||||
{"name": name, "content": content},
|
||||
)
|
||||
|
||||
def upload_pad_attachment(self, pad_slug: str, attachment: SourceAttachment) -> dict[str, Any]:
|
||||
return self.upload_multipart(
|
||||
f"/api/pads/{quote(pad_slug, safe='')}/files",
|
||||
filename=attachment.filename,
|
||||
data=attachment.data,
|
||||
mime_type=attachment.mime_type,
|
||||
)
|
||||
|
||||
def upload_workspace_note_attachment(
|
||||
self,
|
||||
workspace_slug: str,
|
||||
note_slug: str,
|
||||
attachment: SourceAttachment,
|
||||
) -> dict[str, Any]:
|
||||
return self.upload_multipart(
|
||||
(
|
||||
f"/api/workspaces/{quote(workspace_slug, safe='')}/notes/"
|
||||
f"{quote(note_slug, safe='')}/files"
|
||||
),
|
||||
filename=attachment.filename,
|
||||
data=attachment.data,
|
||||
mime_type=attachment.mime_type,
|
||||
)
|
||||
|
||||
|
||||
class ReadableHtmlParser(HTMLParser):
|
||||
BLOCK_TAGS = {"h1", "h2", "h3", "h4", "p", "li", "blockquote", "pre", "br"}
|
||||
SKIP_TAGS = {"script", "style", "svg", "noscript", "nav", "footer"}
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__(convert_charrefs=True)
|
||||
self.parts: list[str] = []
|
||||
self.skip_depth = 0
|
||||
self.title = "Wikipedia snapshot"
|
||||
self.in_title = False
|
||||
self.images: list[tuple[str, str, int | None, int | None]] = []
|
||||
|
||||
@staticmethod
|
||||
def _dimension(value: str | None) -> int | None:
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
return int(float(value))
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _srcset_url(value: str | None) -> str:
|
||||
candidates = []
|
||||
for item in str(value or "").split(","):
|
||||
parts = item.strip().split()
|
||||
if not parts:
|
||||
continue
|
||||
descriptor = parts[1] if len(parts) > 1 else "1x"
|
||||
try:
|
||||
weight = float(descriptor.removesuffix("w").removesuffix("x"))
|
||||
except ValueError:
|
||||
weight = 1.0
|
||||
candidates.append((weight, parts[0]))
|
||||
if not candidates:
|
||||
return ""
|
||||
candidates.sort()
|
||||
for weight, url in candidates:
|
||||
if weight >= 640:
|
||||
return url
|
||||
return candidates[-1][1]
|
||||
|
||||
def handle_starttag(self, tag: str, attrs: list[tuple[str, str | None]]) -> None:
|
||||
attributes = {name.lower(): value for name, value in attrs}
|
||||
if tag in self.SKIP_TAGS:
|
||||
self.skip_depth += 1
|
||||
if tag == "title":
|
||||
self.in_title = True
|
||||
if tag == "img" and not self.skip_depth:
|
||||
src = str(attributes.get("data-src") or attributes.get("src") or "")
|
||||
src = src or self._srcset_url(attributes.get("srcset"))
|
||||
if src:
|
||||
alt = " ".join(str(attributes.get("alt") or "Wikipedia image").split())
|
||||
self.images.append((src, alt, self._dimension(attributes.get("width")), self._dimension(attributes.get("height"))))
|
||||
if not self.skip_depth and tag in self.BLOCK_TAGS:
|
||||
self.parts.append("\n")
|
||||
if tag == "li":
|
||||
self.parts.append("- ")
|
||||
|
||||
def handle_endtag(self, tag: str) -> None:
|
||||
if tag == "title":
|
||||
self.in_title = False
|
||||
if tag in self.SKIP_TAGS and self.skip_depth:
|
||||
self.skip_depth -= 1
|
||||
if not self.skip_depth and tag in self.BLOCK_TAGS:
|
||||
self.parts.append("\n")
|
||||
|
||||
def handle_data(self, data: str) -> None:
|
||||
if self.skip_depth:
|
||||
return
|
||||
text = " ".join(data.split())
|
||||
if not text:
|
||||
return
|
||||
if self.in_title:
|
||||
self.title = text.removesuffix(" - Wikipedia")
|
||||
return
|
||||
self.parts.append(text + " ")
|
||||
|
||||
def wikipedia_images(self, source_url: str, limit: int) -> list[tuple[str, str]]:
|
||||
result: list[tuple[str, str]] = []
|
||||
seen: set[str] = set()
|
||||
for raw_url, alt, width, height in self.images:
|
||||
url = urljoin(source_url, raw_url)
|
||||
parsed = urlparse(url)
|
||||
host = parsed.hostname or ""
|
||||
path = parsed.path.lower()
|
||||
if parsed.scheme != "https" or not host.endswith(WIKIMEDIA_IMAGE_HOST_SUFFIX):
|
||||
continue
|
||||
if not path.endswith(WIKIPEDIA_IMAGE_EXTENSIONS):
|
||||
continue
|
||||
if (width is not None and width < 160) or (height is not None and height < 120):
|
||||
continue
|
||||
if any(fragment in path for fragment in ("/icons/", "wikimedia-button", "poweredby_mediawiki", "commons-logo")):
|
||||
continue
|
||||
if url in seen:
|
||||
continue
|
||||
seen.add(url)
|
||||
safe_alt = alt.replace("[", "(").replace("]", ")").replace("\n", " ").strip()
|
||||
result.append((url, safe_alt or "Wikipedia image"))
|
||||
if len(result) >= limit:
|
||||
break
|
||||
return result
|
||||
|
||||
def markdown(self, source_url: str, attachments: tuple[SourceAttachment, ...]) -> str:
|
||||
lines = [" ".join(line.split()) for line in "".join(self.parts).splitlines()]
|
||||
lines = [line for line in lines if line]
|
||||
content = "\n\n".join(lines[:350])
|
||||
image_markdown = "\n\n".join(
|
||||
f"[image={attachment.filename},{attachment.label}]"
|
||||
for attachment in attachments
|
||||
)
|
||||
return (
|
||||
f"# {self.title}\n\n"
|
||||
f"> Test-data snapshot from Wikipedia. Source: {source_url}\n\n"
|
||||
f"{image_markdown}\n\n"
|
||||
f"{content}\n"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SourceAttachment:
|
||||
filename: str
|
||||
data: bytes
|
||||
mime_type: str
|
||||
label: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SourceDocument:
|
||||
title: str
|
||||
content: str
|
||||
attachments: tuple[SourceAttachment, ...] = ()
|
||||
|
||||
|
||||
class SourcePool:
|
||||
def __init__(self, documents: list[SourceDocument], seed: int | None) -> None:
|
||||
self.documents = documents
|
||||
self.random = random.Random(seed)
|
||||
self.lock = threading.Lock()
|
||||
|
||||
def for_item(self, index: int, label: str) -> SourceDocument:
|
||||
with self.lock:
|
||||
source = self.random.choice(self.documents)
|
||||
suffix = "".join(self.random.choices(string.ascii_lowercase + string.digits, k=8))
|
||||
title = f"{label} {index:06d} {suffix}"
|
||||
content = (
|
||||
f"{source.content.rstrip()}\n\n"
|
||||
f"---\n\nLoad-test item: `{label}-{index:06d}-{suffix}`\n"
|
||||
)
|
||||
encoded = content.encode("utf-8")
|
||||
if len(encoded) > MAX_DOCUMENT_BYTES:
|
||||
content = encoded[:MAX_DOCUMENT_BYTES].decode("utf-8", errors="ignore")
|
||||
return SourceDocument(
|
||||
title=title[:80],
|
||||
content=content,
|
||||
attachments=source.attachments,
|
||||
)
|
||||
|
||||
|
||||
def generated_document(index: int) -> SourceDocument:
|
||||
rng = random.Random(index * 7919 + 17)
|
||||
words = [
|
||||
"architecture", "latency", "workspace", "revision", "markdown", "session",
|
||||
"security", "collaboration", "storage", "deployment", "monitoring", "testing",
|
||||
]
|
||||
paragraphs = []
|
||||
for paragraph_index in range(8):
|
||||
sentence_words = [rng.choice(words) for _ in range(rng.randint(35, 70))]
|
||||
paragraphs.append(" ".join(sentence_words).capitalize() + ".")
|
||||
content = (
|
||||
f"# Generated document {index}\n\n"
|
||||
f"- [ ] Validate record {index}\n"
|
||||
f"- [x] Generate deterministic content\n"
|
||||
f"- [ ] Review WebSocket diagnostics\n\n"
|
||||
+ "\n\n".join(paragraphs)
|
||||
+ f"\n\n```json\n{{\"index\": {index}, \"seed\": {rng.randint(1, 999999)}}}\n```\n"
|
||||
)
|
||||
return SourceDocument(title=f"Generated source {index}", content=content)
|
||||
|
||||
|
||||
def image_extension(url: str, mime_type: str) -> str:
|
||||
mapping = {
|
||||
"image/avif": ".avif",
|
||||
"image/bmp": ".bmp",
|
||||
"image/gif": ".gif",
|
||||
"image/jpeg": ".jpg",
|
||||
"image/png": ".png",
|
||||
"image/webp": ".webp",
|
||||
}
|
||||
if mime_type in mapping:
|
||||
return mapping[mime_type]
|
||||
path_extension = os.path.splitext(urlparse(url).path)[1].lower()
|
||||
if path_extension in WIKIPEDIA_IMAGE_EXTENSIONS:
|
||||
return ".jpg" if path_extension == ".jpeg" else path_extension
|
||||
return mimetypes.guess_extension(mime_type) or ".bin"
|
||||
|
||||
|
||||
def download_wikipedia_attachment(
|
||||
opener: Any,
|
||||
url: str,
|
||||
label: str,
|
||||
index: int,
|
||||
timeout: float,
|
||||
) -> SourceAttachment:
|
||||
request = Request(
|
||||
url,
|
||||
headers={
|
||||
"Accept": "image/avif,image/webp,image/png,image/jpeg,image/gif,image/bmp;q=0.9,*/*;q=0.1",
|
||||
"User-Agent": "RustPad-random-data/1.2 (test data generator)",
|
||||
},
|
||||
)
|
||||
with opener.open(request, timeout=timeout) as response:
|
||||
content_length = response.headers.get("Content-Length")
|
||||
if content_length:
|
||||
try:
|
||||
if int(content_length) > MAX_WIKIPEDIA_ATTACHMENT_BYTES:
|
||||
raise RuntimeError(f"Wikipedia image exceeds {MAX_WIKIPEDIA_ATTACHMENT_BYTES} bytes")
|
||||
except ValueError:
|
||||
pass
|
||||
mime_type = str(response.headers.get_content_type() or "application/octet-stream").lower()
|
||||
if mime_type not in WIKIPEDIA_IMAGE_MIME_TYPES:
|
||||
raise RuntimeError(f"Unsupported Wikipedia image type: {mime_type}")
|
||||
data = response.read(MAX_WIKIPEDIA_ATTACHMENT_BYTES + 1)
|
||||
if len(data) > MAX_WIKIPEDIA_ATTACHMENT_BYTES:
|
||||
raise RuntimeError(f"Wikipedia image exceeds {MAX_WIKIPEDIA_ATTACHMENT_BYTES} bytes")
|
||||
if not data:
|
||||
raise RuntimeError("Wikipedia image is empty")
|
||||
extension = image_extension(url, mime_type)
|
||||
filename = f"wikipedia-{index:02d}{extension}"
|
||||
safe_label = label.replace("]", ")").replace("\r", " ").replace("\n", " ").strip()
|
||||
return SourceAttachment(
|
||||
filename=filename,
|
||||
data=data,
|
||||
mime_type=mime_type,
|
||||
label=safe_label or f"Wikipedia image {index}",
|
||||
)
|
||||
|
||||
|
||||
def fetch_wikipedia_snapshot(
|
||||
index: int,
|
||||
timeout: float,
|
||||
insecure: bool,
|
||||
image_limit: int,
|
||||
attempts: int,
|
||||
) -> SourceDocument:
|
||||
context = ssl._create_unverified_context() if insecure else ssl.create_default_context()
|
||||
opener = build_opener(HTTPSHandler(context=context))
|
||||
last_error: Exception | None = None
|
||||
for attempt in range(1, attempts + 1):
|
||||
request = Request(
|
||||
WIKIPEDIA_RANDOM_URL,
|
||||
headers={"User-Agent": "RustPad-random-data/1.1 (test data generator)"},
|
||||
)
|
||||
try:
|
||||
with opener.open(request, timeout=timeout) as response:
|
||||
raw = response.read(MAX_SOURCE_BYTES)
|
||||
source_url = response.geturl()
|
||||
charset = response.headers.get_content_charset() or "utf-8"
|
||||
parser = ReadableHtmlParser()
|
||||
parser.feed(raw.decode(charset, errors="replace"))
|
||||
image_candidates = parser.wikipedia_images(source_url, image_limit * 3)
|
||||
if not image_candidates:
|
||||
last_error = RuntimeError(f"Wikipedia page had no usable images (attempt {attempt}/{attempts})")
|
||||
continue
|
||||
attachments: list[SourceAttachment] = []
|
||||
for image_url, label in image_candidates:
|
||||
try:
|
||||
attachments.append(
|
||||
download_wikipedia_attachment(
|
||||
opener,
|
||||
image_url,
|
||||
label,
|
||||
len(attachments) + 1,
|
||||
timeout,
|
||||
)
|
||||
)
|
||||
except Exception as error: # noqa: BLE001 - another candidate may still work.
|
||||
last_error = error
|
||||
continue
|
||||
if len(attachments) >= image_limit:
|
||||
break
|
||||
if not attachments:
|
||||
last_error = RuntimeError(
|
||||
f"Wikipedia page images could not be downloaded (attempt {attempt}/{attempts}): {last_error}"
|
||||
)
|
||||
continue
|
||||
attachment_tuple = tuple(attachments)
|
||||
content = parser.markdown(source_url, attachment_tuple)
|
||||
return SourceDocument(
|
||||
title=parser.title or f"Wikipedia {index}",
|
||||
content=content,
|
||||
attachments=attachment_tuple,
|
||||
)
|
||||
except Exception as error: # noqa: BLE001 - retries cover transient Wikipedia failures.
|
||||
last_error = error
|
||||
raise RuntimeError(f"Could not fetch a Wikipedia article with images: {last_error}")
|
||||
|
||||
|
||||
def build_source_pool(args: argparse.Namespace, total_items: int) -> SourcePool:
|
||||
pool_size = max(1, min(args.source_pool_size, max(1, total_items)))
|
||||
if args.source == "generated":
|
||||
return SourcePool([generated_document(index) for index in range(pool_size)], args.seed)
|
||||
|
||||
documents: list[SourceDocument] = []
|
||||
failures = 0
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=min(args.source_workers, pool_size)) as executor:
|
||||
futures = [
|
||||
executor.submit(
|
||||
fetch_wikipedia_snapshot,
|
||||
index,
|
||||
args.timeout,
|
||||
args.insecure,
|
||||
args.wikipedia_images,
|
||||
args.wikipedia_attempts,
|
||||
)
|
||||
for index in range(pool_size)
|
||||
]
|
||||
for index, future in enumerate(concurrent.futures.as_completed(futures), start=1):
|
||||
try:
|
||||
documents.append(future.result())
|
||||
print(f"Source snapshots: {index}/{pool_size}", end="\r", flush=True)
|
||||
except Exception as error: # noqa: BLE001 - all source failures are reported below.
|
||||
failures += 1
|
||||
print(f"\nWikipedia source failed: {error}", file=sys.stderr)
|
||||
print()
|
||||
if not documents:
|
||||
raise RuntimeError("No Wikipedia snapshot with a usable image could be downloaded.")
|
||||
if failures:
|
||||
print(
|
||||
f"Using {len(documents)} Wikipedia snapshots; {failures} source downloads failed. "
|
||||
"No generated fallback was added.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return SourcePool(documents, args.seed)
|
||||
|
||||
|
||||
class Progress:
|
||||
def __init__(self, total: int) -> None:
|
||||
self.total = total
|
||||
self.completed = 0
|
||||
self.failed = 0
|
||||
self.started = time.monotonic()
|
||||
self.lock = threading.Lock()
|
||||
|
||||
def record(self, success: bool) -> None:
|
||||
with self.lock:
|
||||
self.completed += 1
|
||||
if not success:
|
||||
self.failed += 1
|
||||
if self.completed == self.total or self.completed % 100 == 0:
|
||||
elapsed = max(0.001, time.monotonic() - self.started)
|
||||
rate = self.completed / elapsed
|
||||
print(
|
||||
f"Created: {self.completed}/{self.total} | failures: {self.failed} | {rate:.1f}/s",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
|
||||
def execute_tasks(
|
||||
tasks: Iterable[tuple[str, str, str | None, int]],
|
||||
*,
|
||||
client: RustPadClient,
|
||||
sources: SourcePool,
|
||||
workers: int,
|
||||
total: int,
|
||||
) -> tuple[list[str], int]:
|
||||
progress = Progress(total)
|
||||
failures: list[str] = []
|
||||
failures_lock = threading.Lock()
|
||||
|
||||
def run(task: tuple[str, str, str | None, int]) -> None:
|
||||
kind, label, workspace_slug, index = task
|
||||
document = sources.for_item(index, label)
|
||||
try:
|
||||
if kind == "pad":
|
||||
created = client.create_pad(document.title, document.content)
|
||||
pad_slug = str(created.get("slug") or "")
|
||||
if not pad_slug:
|
||||
raise RuntimeError("Created pad response did not contain a slug.")
|
||||
for attachment in document.attachments:
|
||||
client.upload_pad_attachment(pad_slug, attachment)
|
||||
else:
|
||||
if not workspace_slug:
|
||||
raise RuntimeError("Workspace slug is missing.")
|
||||
created = client.create_workspace_note(workspace_slug, document.title, document.content)
|
||||
note_slug = str(created.get("slug") or "")
|
||||
if not note_slug:
|
||||
raise RuntimeError("Created workspace note response did not contain a slug.")
|
||||
for attachment in document.attachments:
|
||||
client.upload_workspace_note_attachment(
|
||||
workspace_slug,
|
||||
note_slug,
|
||||
attachment,
|
||||
)
|
||||
progress.record(True)
|
||||
except Exception as error: # noqa: BLE001 - all failures are reported after the run.
|
||||
with failures_lock:
|
||||
if len(failures) < 50:
|
||||
failures.append(str(error))
|
||||
progress.record(False)
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=workers) as executor:
|
||||
list(executor.map(run, tasks, chunksize=1))
|
||||
return failures, progress.failed
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Create standalone notes, workspaces and workspace notes through the RustPad API.",
|
||||
)
|
||||
parser.add_argument("--ip", default="localhost", help="RustPad host name or IP address")
|
||||
parser.add_argument("--port", type=int, default=3000, help="RustPad port")
|
||||
parser.add_argument("--scheme", choices=("http", "https"), default="http")
|
||||
parser.add_argument("--base-url", help="Complete base URL; overrides --ip, --port and --scheme")
|
||||
parser.add_argument("--source", choices=("generated", "wikipedia"), default="generated")
|
||||
parser.add_argument("--notes", type=int, default=0, help="Total number of notes to create")
|
||||
parser.add_argument("--workspaces", type=int, default=0, help="Number of workspaces")
|
||||
parser.add_argument(
|
||||
"--notes-in-workspaces",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Number of notes from --notes distributed across all workspaces",
|
||||
)
|
||||
parser.add_argument("--user", required=True, help="Login e-mail or LDAP/AD username")
|
||||
parser.add_argument("--password", help="Account password; otherwise RUSTPAD_TEST_PASSWORD or a prompt is used")
|
||||
parser.add_argument("--workers", type=int, default=12, help="Concurrent API requests")
|
||||
parser.add_argument("--source-workers", type=int, default=6, help="Concurrent website downloads")
|
||||
parser.add_argument("--source-pool-size", type=int, default=40, help="Website/generated source documents reused by test items")
|
||||
parser.add_argument(
|
||||
"--wikipedia-images",
|
||||
type=int,
|
||||
default=3,
|
||||
help="Maximum Wikipedia images downloaded and uploaded as attachments to each note",
|
||||
)
|
||||
parser.add_argument("--wikipedia-attempts", type=int, default=8, help="Random articles tried when a Wikipedia page has no usable image")
|
||||
parser.add_argument("--timeout", type=float, default=30.0, help="HTTP timeout in seconds")
|
||||
parser.add_argument("--retries", type=int, default=5, help="Retries for transient API errors")
|
||||
parser.add_argument("--seed", type=int, help="Deterministic random seed")
|
||||
parser.add_argument("--prefix", default="Load test", help="Workspace name prefix")
|
||||
parser.add_argument("--insecure", action="store_true", help="Disable TLS certificate verification")
|
||||
parser.add_argument("--dry-run", action="store_true", help="Print counts without writing data")
|
||||
args = parser.parse_args()
|
||||
|
||||
for name in ("notes", "workspaces", "notes_in_workspaces"):
|
||||
if getattr(args, name) < 0:
|
||||
parser.error(f"--{name.replace('_', '-')} cannot be negative")
|
||||
if args.workers < 1 or args.source_workers < 1 or args.source_pool_size < 1:
|
||||
parser.error("worker and source-pool values must be positive")
|
||||
if not 1 <= args.wikipedia_images <= 10:
|
||||
parser.error("--wikipedia-images must be between 1 and 10")
|
||||
if args.wikipedia_attempts < 1:
|
||||
parser.error("--wikipedia-attempts must be positive")
|
||||
if args.notes_in_workspaces > args.notes:
|
||||
parser.error("--notes-in-workspaces cannot exceed --notes")
|
||||
if args.notes_in_workspaces and not args.workspaces:
|
||||
parser.error("--workspaces must be positive when --notes-in-workspaces is used")
|
||||
return args
|
||||
|
||||
|
||||
def base_url(args: argparse.Namespace) -> str:
|
||||
if args.base_url:
|
||||
return args.base_url.rstrip("/")
|
||||
host = f"[{args.ip}]" if ":" in args.ip and not args.ip.startswith("[") else args.ip
|
||||
return f"{args.scheme}://{host}:{args.port}"
|
||||
|
||||
|
||||
def resolve_password(args: argparse.Namespace) -> str:
|
||||
password = args.password or os.environ.get("RUSTPAD_TEST_PASSWORD")
|
||||
if password:
|
||||
return password
|
||||
if not sys.stdin.isatty():
|
||||
raise RuntimeError("Set --password or RUSTPAD_TEST_PASSWORD when stdin is not interactive.")
|
||||
return getpass.getpass("RustPad password: ")
|
||||
|
||||
|
||||
def main() -> int:
|
||||
args = parse_args()
|
||||
random.seed(args.seed)
|
||||
workspace_note_total = args.notes_in_workspaces
|
||||
standalone_note_total = args.notes - workspace_note_total
|
||||
item_total = args.notes
|
||||
per_workspace = []
|
||||
if args.workspaces:
|
||||
base_count, remainder = divmod(workspace_note_total, args.workspaces)
|
||||
per_workspace = [
|
||||
base_count + (1 if index < remainder else 0)
|
||||
for index in range(args.workspaces)
|
||||
]
|
||||
print(
|
||||
f"Target: {args.notes} notes total: {standalone_note_total} standalone and "
|
||||
f"{workspace_note_total} across {args.workspaces} workspaces."
|
||||
)
|
||||
if per_workspace:
|
||||
minimum = min(per_workspace)
|
||||
maximum = max(per_workspace)
|
||||
distribution = str(minimum) if minimum == maximum else f"{minimum}-{maximum}"
|
||||
print(f"Workspace distribution: {distribution} notes per workspace.")
|
||||
if args.dry_run:
|
||||
return 0
|
||||
if item_total == 0 and args.workspaces == 0:
|
||||
print("Nothing to create.")
|
||||
return 0
|
||||
|
||||
client = RustPadClient(
|
||||
base_url(args),
|
||||
timeout=args.timeout,
|
||||
retries=args.retries,
|
||||
insecure=args.insecure,
|
||||
)
|
||||
session = client.login(args.user, resolve_password(args))
|
||||
print(f"Logged in as {session.get('nickname') or args.user}.")
|
||||
|
||||
sources = build_source_pool(args, max(1, item_total))
|
||||
|
||||
workspace_slugs: list[str] = []
|
||||
for index in range(1, args.workspaces + 1):
|
||||
suffix = "".join(random.choices(string.ascii_lowercase + string.digits, k=6))
|
||||
name = f"{args.prefix} workspace {index:04d} {suffix}"[:80]
|
||||
workspace = client.create_workspace(name)
|
||||
workspace_slugs.append(str(workspace["slug"]))
|
||||
print(f"Workspaces: {index}/{args.workspaces}", end="\r", flush=True)
|
||||
if args.workspaces:
|
||||
print()
|
||||
|
||||
tasks: list[tuple[str, str, str | None, int]] = []
|
||||
for index in range(1, standalone_note_total + 1):
|
||||
tasks.append(("pad", f"Standalone note", None, index))
|
||||
absolute_index = standalone_note_total
|
||||
for workspace_index, (slug, note_count) in enumerate(
|
||||
zip(workspace_slugs, per_workspace, strict=True),
|
||||
start=1,
|
||||
):
|
||||
for _note_index in range(1, note_count + 1):
|
||||
absolute_index += 1
|
||||
tasks.append(("workspace-note", f"Workspace {workspace_index:04d} note", slug, absolute_index))
|
||||
|
||||
failures, failed_count = execute_tasks(
|
||||
tasks,
|
||||
client=client,
|
||||
sources=sources,
|
||||
workers=args.workers,
|
||||
total=len(tasks),
|
||||
) if tasks else ([], 0)
|
||||
|
||||
if failed_count:
|
||||
print(f"Completed with {failed_count} failures ({len(failures)} shown):", file=sys.stderr)
|
||||
for failure in failures:
|
||||
print(f"- {failure}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
print("Data generation completed successfully.")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user