Merge branch 'main' into feature/chat-api
This commit is contained in:
commit
e43bb2b2f4
5 changed files with 203 additions and 21 deletions
|
|
@ -48,13 +48,31 @@ class ScrapeConfig:
|
|||
max_comments_per_item: int
|
||||
|
||||
|
||||
def _resolve_token(token: str) -> str:
|
||||
tok = token or os.environ.get("GH_TOKEN", "") or os.environ.get("GITHUB_TOKEN", "")
|
||||
if not tok:
|
||||
raise ValueError(
|
||||
"GitHub token is required. Set it in the recipe config or the GH_TOKEN / GITHUB_TOKEN env var."
|
||||
@dataclass(frozen = True)
|
||||
class ResolvedToken:
|
||||
value: str
|
||||
source: str
|
||||
|
||||
|
||||
def _resolve_token(token: str) -> ResolvedToken:
|
||||
if token:
|
||||
return ResolvedToken(
|
||||
value = token,
|
||||
source = "explicit token argument (recipe-level field)",
|
||||
)
|
||||
return tok
|
||||
if os.environ.get("GH_TOKEN"):
|
||||
return ResolvedToken(
|
||||
value = os.environ["GH_TOKEN"],
|
||||
source = "GH_TOKEN environment variable",
|
||||
)
|
||||
if os.environ.get("GITHUB_TOKEN"):
|
||||
return ResolvedToken(
|
||||
value = os.environ["GITHUB_TOKEN"],
|
||||
source = "GITHUB_TOKEN environment variable",
|
||||
)
|
||||
raise ValueError(
|
||||
"GitHub token is required. Set it in the recipe config or the GH_TOKEN / GITHUB_TOKEN env var."
|
||||
)
|
||||
|
||||
|
||||
def _read_jsonl(path: Path, max_rows: int | None = None):
|
||||
|
|
@ -155,7 +173,7 @@ def _flatten_commit_row(r: dict, repo: str) -> dict:
|
|||
def scrape(cfg: ScrapeConfig, base_dir: Path):
|
||||
token = _resolve_token(cfg.token)
|
||||
GitHubClient, RepoScraper = _load_impl()
|
||||
client = GitHubClient(token = token)
|
||||
client = GitHubClient(token = token.value, token_source = token.source)
|
||||
base_dir.mkdir(parents = True, exist_ok = True)
|
||||
|
||||
# Per-resource trial limits. limit <= 0 means "all": use a very large cap.
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ import json
|
|||
import os
|
||||
import time
|
||||
import logging
|
||||
from datetime import timezone
|
||||
from email.utils import parsedate_to_datetime
|
||||
from typing import Any, Dict, Iterable, Iterator, List, Optional
|
||||
|
||||
import requests
|
||||
|
|
@ -29,16 +31,46 @@ class RateLimitError(Exception):
|
|||
pass
|
||||
|
||||
|
||||
class GitHubAuthError(RuntimeError):
|
||||
"""Raised when GitHub returns 401/403 due to invalid or insufficient credentials."""
|
||||
|
||||
|
||||
def _retry_after_seconds(value: str | None) -> int | None:
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
return max(0, int(value))
|
||||
except ValueError:
|
||||
pass
|
||||
try:
|
||||
retry_at = parsedate_to_datetime(value)
|
||||
except (TypeError, ValueError, IndexError, OverflowError):
|
||||
return None
|
||||
if retry_at.tzinfo is None:
|
||||
retry_at = retry_at.replace(tzinfo = timezone.utc)
|
||||
return max(0, int(retry_at.timestamp() - time.time()))
|
||||
|
||||
|
||||
class GitHubClient:
|
||||
def __init__(
|
||||
self,
|
||||
min_remaining_graphql: int = 100,
|
||||
min_remaining_rest: int = 100,
|
||||
token: str | None = None,
|
||||
token_source: str | None = None,
|
||||
):
|
||||
token = token or os.environ.get("GH_TOKEN") or os.environ.get("GITHUB_TOKEN")
|
||||
if not token:
|
||||
raise RuntimeError("GH_TOKEN not set in environment")
|
||||
if token:
|
||||
self._token_source = (
|
||||
token_source or "explicit token argument (recipe-level field)"
|
||||
)
|
||||
elif os.environ.get("GH_TOKEN"):
|
||||
self._token_source = "GH_TOKEN environment variable"
|
||||
token = os.environ["GH_TOKEN"]
|
||||
elif os.environ.get("GITHUB_TOKEN"):
|
||||
self._token_source = "GITHUB_TOKEN environment variable"
|
||||
token = os.environ["GITHUB_TOKEN"]
|
||||
else:
|
||||
raise RuntimeError("GH_TOKEN or GITHUB_TOKEN not set in environment")
|
||||
self.session = requests.Session()
|
||||
self.session.headers.update(
|
||||
{**BASE_HEADERS, "Authorization": f"Bearer {token}"}
|
||||
|
|
@ -59,6 +91,49 @@ class GitHubClient:
|
|||
log.warning("Rate limit hit. Sleeping %ds until reset.", wait)
|
||||
time.sleep(wait)
|
||||
|
||||
def _is_rate_limit_response(self, r: "requests.Response") -> bool:
|
||||
if r.headers.get("Retry-After"):
|
||||
return True
|
||||
if r.headers.get("X-RateLimit-Remaining") == "0":
|
||||
return True
|
||||
body = (r.text or "").lower()
|
||||
return any(
|
||||
marker in body
|
||||
for marker in (
|
||||
"api rate limit exceeded",
|
||||
"rate limit exceeded",
|
||||
"secondary rate limit",
|
||||
"secondary limit",
|
||||
"abuse detection mechanism",
|
||||
"abuse detection",
|
||||
)
|
||||
)
|
||||
|
||||
def _is_auth_failure(self, r: "requests.Response") -> bool:
|
||||
"""Distinguish auth failures from rate limiting on 401/403 responses.
|
||||
|
||||
- 401: always an auth failure (invalid / expired / wrong-scope token).
|
||||
- 403: an auth failure UNLESS the response carries a clear rate-limit signal
|
||||
(Retry-After header, X-RateLimit-Remaining: 0, or GitHub's secondary /
|
||||
abuse rate-limit response text).
|
||||
"""
|
||||
if r.status_code == 401:
|
||||
return True
|
||||
if r.status_code == 403:
|
||||
return not self._is_rate_limit_response(r)
|
||||
return False
|
||||
|
||||
def _raise_auth_error(self, r: "requests.Response", endpoint: str) -> None:
|
||||
snippet = (r.text or "").strip()[:200]
|
||||
request_id = r.headers.get("X-GitHub-Request-Id")
|
||||
request_id_message = f" Request ID: {request_id}." if request_id else ""
|
||||
raise GitHubAuthError(
|
||||
f"GitHub {endpoint} returned {r.status_code} {r.reason}. "
|
||||
f"Token source: {self._token_source}. "
|
||||
f"The token is invalid, expired, or missing required scopes — "
|
||||
f"retrying will not recover.{request_id_message} Response: {snippet}"
|
||||
)
|
||||
|
||||
def _check_rate_and_wait(self, kind: str) -> None:
|
||||
if kind == "graphql":
|
||||
remaining = self.graphql_remaining
|
||||
|
|
@ -112,13 +187,14 @@ class GitHubClient:
|
|||
time.sleep(backoff)
|
||||
backoff = min(backoff * 2, 60)
|
||||
continue
|
||||
if self._is_auth_failure(r):
|
||||
self._raise_auth_error(r, "GraphQL")
|
||||
if r.status_code == 403 or r.status_code == 429:
|
||||
# Check for secondary/abuse
|
||||
retry_after = r.headers.get("Retry-After")
|
||||
if retry_after:
|
||||
t = int(retry_after)
|
||||
log.warning("Secondary rate limit. Sleep %ds.", t)
|
||||
time.sleep(t + 2)
|
||||
retry_after = _retry_after_seconds(r.headers.get("Retry-After"))
|
||||
if retry_after is not None:
|
||||
log.warning("Secondary rate limit. Sleep %ds.", retry_after)
|
||||
time.sleep(retry_after + 2)
|
||||
continue
|
||||
if self.graphql_reset:
|
||||
self._sleep_until(self.graphql_reset)
|
||||
|
|
@ -188,12 +264,15 @@ class GitHubClient:
|
|||
time.sleep(backoff)
|
||||
backoff = min(backoff * 2, 60)
|
||||
continue
|
||||
if self._is_auth_failure(r):
|
||||
self._raise_auth_error(r, "REST")
|
||||
if r.status_code in (403, 429):
|
||||
retry_after = r.headers.get("Retry-After")
|
||||
if retry_after:
|
||||
t = int(retry_after)
|
||||
log.warning("Secondary rate limit on REST. Sleep %ds.", t)
|
||||
time.sleep(t + 2)
|
||||
retry_after = _retry_after_seconds(r.headers.get("Retry-After"))
|
||||
if retry_after is not None:
|
||||
log.warning(
|
||||
"Secondary rate limit on REST. Sleep %ds.", retry_after
|
||||
)
|
||||
time.sleep(retry_after + 2)
|
||||
continue
|
||||
# Check if primary rate
|
||||
if self.rest_remaining == 0 and self.rest_reset:
|
||||
|
|
|
|||
|
|
@ -42,7 +42,9 @@ anyio
|
|||
sniffio
|
||||
h11
|
||||
|
||||
tokenizers
|
||||
# Unpinned resolves to 0.23.1+ which breaks `from transformers import
|
||||
# AutoConfig`; transformers 4.56..5.3 declares tokenizers<=0.23.0.
|
||||
tokenizers<=0.23.0
|
||||
transformers>=4.51.3,!=4.52.0,!=4.52.1,!=4.52.2,!=4.52.3,!=4.53.0,!=4.54.0,!=4.55.0,!=4.55.1,!=4.57.0,!=4.57.4,!=4.57.5,!=5.0.0,!=5.1.0,<=5.3.0
|
||||
trl>=0.18.2,!=0.19.0,<=0.24.0
|
||||
sentence-transformers
|
||||
|
|
|
|||
69
tests/python/test_patch_trl_rl_trainers_defensive.py
Normal file
69
tests/python/test_patch_trl_rl_trainers_defensive.py
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
|
||||
|
||||
"""Regression tests: _patch_trl_rl_trainers must never raise.
|
||||
|
||||
The wrapper in unsloth/models/rl.py ring-fences the impl so direct
|
||||
callers (CI shims, downstream tools) don't have to. Lock that
|
||||
contract here.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
pytest.importorskip("trl")
|
||||
|
||||
|
||||
def _import_helpers():
|
||||
try:
|
||||
from unsloth.models.rl import (
|
||||
_patch_trl_rl_trainers,
|
||||
_patch_trl_rl_trainers_impl,
|
||||
)
|
||||
except ImportError as e:
|
||||
pytest.skip(f"unsloth.models.rl helpers not importable: {e}")
|
||||
return _patch_trl_rl_trainers, _patch_trl_rl_trainers_impl
|
||||
|
||||
|
||||
def test_patch_trl_rl_trainers_swallows_unknown_trainer_name():
|
||||
wrapper, _impl = _import_helpers()
|
||||
assert wrapper("definitely_not_a_real_trainer_xyz") is None
|
||||
|
||||
|
||||
def test_patch_trl_rl_trainers_swallows_garbage_input():
|
||||
wrapper, _impl = _import_helpers()
|
||||
for bad in ("", "..", "trainer with space", "sft_trainer; rm -rf /"):
|
||||
assert wrapper(bad) is None, f"raised on input: {bad!r}"
|
||||
|
||||
|
||||
def test_impl_is_separately_exposed():
|
||||
# Power users can still call the impl directly for the raising path.
|
||||
_wrapper, impl = _import_helpers()
|
||||
assert callable(impl)
|
||||
|
||||
|
||||
def test_wrapper_delegates_to_impl(monkeypatch):
|
||||
from unsloth.models import rl as _rl
|
||||
|
||||
sentinel = object()
|
||||
calls = []
|
||||
|
||||
def _fake_impl(trainer_file):
|
||||
calls.append(trainer_file)
|
||||
return sentinel
|
||||
|
||||
monkeypatch.setattr(_rl, "_patch_trl_rl_trainers_impl", _fake_impl)
|
||||
assert _rl._patch_trl_rl_trainers("sft_trainer") is sentinel
|
||||
assert calls == ["sft_trainer"]
|
||||
|
||||
|
||||
def test_wrapper_swallows_impl_exception(monkeypatch):
|
||||
from unsloth.models import rl as _rl
|
||||
|
||||
def _boom(_trainer_file):
|
||||
raise RuntimeError("simulated TRL 1.x rename failure")
|
||||
|
||||
monkeypatch.setattr(_rl, "_patch_trl_rl_trainers_impl", _boom)
|
||||
assert _rl._patch_trl_rl_trainers("sft_trainer") is None
|
||||
|
|
@ -540,6 +540,20 @@ def _wrap_grpo_generate_and_score(trainer_cls):
|
|||
|
||||
|
||||
def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
||||
# Defensive wrapper: matches patch_trl_rl_trainers()'s try/except so
|
||||
# direct callers don't see exceptions from the impl on TRL versions
|
||||
# that rename or move classes (e.g. TRL 1.x trl.experimental).
|
||||
try:
|
||||
return _patch_trl_rl_trainers_impl(trainer_file)
|
||||
except Exception as e:
|
||||
logger.info(
|
||||
f"Unsloth: Could not patch trl.trainer.{trainer_file}: "
|
||||
f"{type(e).__name__}: {e}"
|
||||
)
|
||||
return
|
||||
|
||||
|
||||
def _patch_trl_rl_trainers_impl(trainer_file = "grpo_trainer"):
|
||||
# Patch for vLLM and Unsloth PEFT
|
||||
import trl
|
||||
import trl.trainer
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue