Skip to content
1 change: 1 addition & 0 deletions changelog/1340.fixed.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
When an API key is valid but the model license has not been accepted, TabPFN now opens the license acceptance page directly, or links to it when no browser is available.
211 changes: 161 additions & 50 deletions src/tabpfn/browser_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,8 +22,10 @@
import urllib.parse
import urllib.request
import webbrowser
from enum import Enum
from pathlib import Path
from typing import TYPE_CHECKING
from typing_extensions import assert_never

from tabpfn.errors import (
TabPFNError,
Expand All @@ -42,6 +44,43 @@
_accepted_repos: set[str] = set()


class _AuthStep(Enum):
"""What the user still has to do in the browser."""

# No usable API key: log in (or register), then accept the license.
LOGIN = "login"
# The API key is valid; only the license acceptance is missing.
ACCEPT_LICENSE = "accept_license"


class _NoBrowserToken(Enum):
"""Why the browser flow produced no API key."""

# TABPFN_NO_BROWSER is set, so the browser flow never ran.
DISABLED = "disabled"
# No interactive terminal, or the user aborted the flow.
UNAVAILABLE = "unavailable"


def _browser_url(
gui_url: str,
step: _AuthStep,
*,
hf_repo_id: str,
callback_url: str | None = None,
) -> str:
"""Return the web app page that completes ``step``."""
match step:
case _AuthStep.LOGIN:
path = "login"
case _AuthStep.ACCEPT_LICENSE:
path = "accept-license"
case _:
assert_never(step)
url = f"{gui_url}/{path}?hf_repo_id={urllib.parse.quote(hf_repo_id)}"
return f"{url}&callback={callback_url}" if callback_url else url


# ---------------------------------------------------------------------------
# Environment detection
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -353,19 +392,15 @@ def _poll_for_token(
return received_token[0]


def _headless_interactive_login(
gui_url: str, hf_repo_id: str | None = None
) -> str | None:
def _headless_interactive_login(gui_url: str, *, hf_repo_id: str) -> str | None:
"""Token acquisition for headless but interactive environments (e.g. SSH).

Shows the login URL, offers single-keypress clipboard copy via OSC 52,
and waits for the user to paste a token.

Returns the JWT on success, or ``None`` on abort / EOF.
"""
login_url = f"{gui_url}/login"
if hf_repo_id:
login_url += f"?hf_repo_id={urllib.parse.quote(hf_repo_id)}"
login_url = _browser_url(gui_url, _AuthStep.LOGIN, hf_repo_id=hf_repo_id)

print( # noqa: T201
"\nTabPFN requires a one-time license acceptance to download"
Expand Down Expand Up @@ -479,7 +514,12 @@ def _headless_readline_loop(login_url: str) -> str | None:
return None


def try_browser_login(gui_url: str, hf_repo_id: str | None = None) -> str | None:
def try_browser_login(
gui_url: str,
*,
hf_repo_id: str,
step: _AuthStep,
) -> str | None:
"""Obtain a token via browser callback and/or manual paste concurrently.

Chooses the right strategy based on the environment:
Expand All @@ -496,7 +536,15 @@ def try_browser_login(gui_url: str, hf_repo_id: str | None = None) -> str | None
return None

if not _has_display():
return _headless_interactive_login(gui_url, hf_repo_id=hf_repo_id)
match step:
case _AuthStep.LOGIN:
return _headless_interactive_login(gui_url, hf_repo_id=hf_repo_id)
case _AuthStep.ACCEPT_LICENSE:
# The key is already set, so there is nothing to paste; the
# caller's error links to the acceptance page instead.
return None
case _:
assert_never(step)

auth_event = threading.Event()
received_token: list[str | None] = [None]
Expand All @@ -508,10 +556,12 @@ def try_browser_login(gui_url: str, hf_repo_id: str | None = None) -> str | None
logger.debug("Could not create callback server", exc_info=True)
return None

callback_url = f"http://localhost:{port}"
login_url = f"{gui_url}/login?callback={callback_url}"
if hf_repo_id:
login_url += f"&hf_repo_id={urllib.parse.quote(hf_repo_id)}"
login_url = _browser_url(
gui_url,
step,
hf_repo_id=hf_repo_id,
callback_url=f"http://localhost:{port}",
)

server_thread = threading.Thread(
target=_serve_until_event, args=(httpd, auth_event), daemon=True
Expand All @@ -522,19 +572,34 @@ def try_browser_login(gui_url: str, hf_repo_id: str | None = None) -> str | None
webbrowser.open(login_url)

# --- print unified instructions ---
match step:
case _AuthStep.LOGIN:
instructions = (
"\nOpening your browser to complete login/registration…\n"
f"\n {login_url}\n"
"\nWaiting for login to complete…\n"
"\nHaving trouble? You can also authenticate manually:\n"
f" 1. Open {gui_url}/account in a browser"
" (log in or register if needed)\n"
f" 2. Accept the license at {gui_url}/account/licenses\n"
" 3. Copy your API Key\n"
" 4. Paste the API key below\n"
)
case _AuthStep.ACCEPT_LICENSE:
instructions = (
"\nYour API key is valid; only the license acceptance is missing."
"\nOpening your browser to accept it…\n"
f"\n {login_url}\n"
"\nWaiting for acceptance to complete…\n"
"\nHaving trouble? Accept the license at the link above,"
" then paste your API key below.\n"
)
case _:
assert_never(step)
print( # noqa: T201
"\nTabPFN requires a one-time license acceptance to download"
" model weights for local inference."
"\nOpening your browser to complete login/registration…\n"
f"\n {login_url}\n"
"\nWaiting for login to complete…\n"
"\nHaving trouble? You can also authenticate manually:\n"
f" 1. Open {gui_url}/account in a browser"
" (log in or register if needed)\n"
" 2. Accept the license at"
f" {gui_url}/account/licenses\n"
" 3. Copy your API Key\n"
" 4. Paste the API key below\n"
f"{instructions}"
)

# --- main thread: poll stdin while waiting for callback ---
Expand All @@ -552,7 +617,7 @@ def try_browser_login(gui_url: str, hf_repo_id: str | None = None) -> str | None
# ---------------------------------------------------------------------------


def ensure_license_accepted(hf_repo_id: str) -> Literal[True]: # noqa: C901
def ensure_license_accepted(hf_repo_id: str) -> Literal[True]: # noqa: C901, PLR0912
"""Ensure the user has accepted the TabPFN license.

Checks for a cached token, verifies it, and falls back to browser login
Expand All @@ -574,6 +639,7 @@ def ensure_license_accepted(hf_repo_id: str) -> Literal[True]: # noqa: C901
api_url = settings.tabpfn.auth_api_url

license_version = _get_license_name(hf_repo_id)
step = _AuthStep.LOGIN

token = get_cached_token()
if token is not None:
Expand All @@ -590,8 +656,8 @@ def ensure_license_accepted(hf_repo_id: str) -> Literal[True]: # noqa: C901
"Could not reach the license server to verify acceptance.\n\n"
"Please check your internet connection and try again."
)
# license_status is False — license not yet accepted.
# Fall through to browser login so the GUI can show the acceptance form.
# license_status is False: the key works, only acceptance is missing.
step = _AuthStep.ACCEPT_LICENSE
logger.info(
"Token valid but license not accepted; opening browser for acceptance.",
)
Expand All @@ -605,31 +671,45 @@ def ensure_license_accepted(hf_repo_id: str) -> Literal[True]: # noqa: C901
logger.info("Cached token is invalid; deleting and re-authenticating.")
delete_cached_token()

# No valid cached token — need browser login.
no_browser = os.environ.get("TABPFN_NO_BROWSER", "").strip()
browser_result: str | _NoBrowserToken
if no_browser and no_browser not in ("0", "false", "no", "off"):
raise TabPFNLicenseError(
"TabPFN requires a one-time license acceptance to download\n"
"model weights for local inference, but browser login is\n"
"disabled (TABPFN_NO_BROWSER is set).\n\n"
"Set the TABPFN_TOKEN environment variable with a valid API key\n"
"obtained from https://ux.priorlabs.ai"
browser_result = _NoBrowserToken.DISABLED
else:
browser_result = (
try_browser_login(gui_url, hf_repo_id=hf_repo_id, step=step)
or _NoBrowserToken.UNAVAILABLE
)

token = try_browser_login(gui_url, hf_repo_id=hf_repo_id)
if token is None:
raise TabPFNLicenseError(
"TabPFN requires a one-time license acceptance to download\n"
"model weights for local inference, but no interactive terminal\n"
"is available.\n\n"
"To authenticate in a non-interactive environment:\n"
f" 1. Open {gui_url} in a browser and log in (or register)\n"
f" 2. Accept the license on the Licenses tab\n"
f" 3. Copy your API Key from {gui_url}/account\n"
' 4. Set the environment variable: export TABPFN_TOKEN="<your-api-key>"\n'
" or in Python (before calling .fit()):"
' import os; os.environ["TABPFN_TOKEN"] = "<your-api-key>"'
)
match step, browser_result:
case _, str() as token:
pass
case _AuthStep.ACCEPT_LICENSE, _NoBrowserToken():
raise _license_not_accepted_error(gui_url, hf_repo_id)
case _AuthStep.LOGIN, _NoBrowserToken.DISABLED:
raise TabPFNLicenseError(
"TabPFN requires a one-time license acceptance to download\n"
"model weights for local inference, but browser login is\n"
"disabled (TABPFN_NO_BROWSER is set).\n\n"
"Set the TABPFN_TOKEN environment variable with a valid API key\n"
"obtained from https://ux.priorlabs.ai"
)
case _AuthStep.LOGIN, _NoBrowserToken.UNAVAILABLE:
raise TabPFNLicenseError(
"TabPFN requires a one-time license acceptance to download\n"
"model weights for local inference, but no interactive terminal\n"
"is available.\n\n"
"To authenticate in a non-interactive environment:\n"
f" 1. Open {gui_url} in a browser and log in (or register)\n"
f" 2. Accept the license on the Licenses tab\n"
f" 3. Copy your API Key from {gui_url}/account\n"
" 4. Set the environment variable:"
' export TABPFN_TOKEN="<your-api-key>"\n'
" or in Python (before calling .fit()):"
' import os; os.environ["TABPFN_TOKEN"] = "<your-api-key>"'
)
case _:
assert_never((step, browser_result))

# Verify the token we just received from the browser.
status = verify_token(token, api_url)
Expand All @@ -644,6 +724,20 @@ def ensure_license_accepted(hf_repo_id: str) -> Literal[True]: # noqa: C901

license_status = check_license_accepted(token, api_url, license_version)
if license_status is True:
# The browser may be signed in to another account than the valid key
# we started with. A TABPFN_TOKEN key outranks the one just saved, so
# without this check every later run would reopen the browser. Only
# the acceptance step had a verified key; after a login, a leftover
# key is just invalid, not another account.
next_run_token = get_cached_token()
if (
step is _AuthStep.ACCEPT_LICENSE
and next_run_token is not None
and next_run_token != token
and check_license_accepted(next_run_token, api_url, license_version)
is False
):
raise _different_account_error(gui_url, hf_repo_id)
Comment thread
cursor[bot] marked this conversation as resolved.
print("License accepted — API key cached for future sessions.\n") # noqa: T201
_accepted_repos.add(hf_repo_id)
return True
Expand All @@ -653,8 +747,25 @@ def ensure_license_accepted(hf_repo_id: str) -> Literal[True]: # noqa: C901
"Please check your internet connection and try again."
)
# license_status is False
encoded = urllib.parse.quote(hf_repo_id)
raise TabPFNLicenseError(
"License not yet accepted. Please complete the acceptance form at\n"
f"{gui_url}/accept-license?hf_repo_id={encoded} and try again."
raise _license_not_accepted_error(gui_url, hf_repo_id)


def _license_not_accepted_error(gui_url: str, hf_repo_id: str) -> TabPFNLicenseError:
accept_url = _browser_url(gui_url, _AuthStep.ACCEPT_LICENSE, hf_repo_id=hf_repo_id)
return TabPFNLicenseError(
f"You are logged in, but the license for {hf_repo_id} has not been\n"
"accepted for your account yet. Accept it once at\n\n"
f" {accept_url}\n\n"
"then try again."
)


def _different_account_error(gui_url: str, hf_repo_id: str) -> TabPFNLicenseError:
accept_url = _browser_url(gui_url, _AuthStep.ACCEPT_LICENSE, hf_repo_id=hf_repo_id)
return TabPFNLicenseError(
f"The license for {hf_repo_id} was accepted in the browser for a\n"
"different account than the one your TABPFN_TOKEN key belongs to.\n"
"Sign in to the browser with that key's account and accept it at\n\n"
f" {accept_url}\n\n"
"then try again."
)
Loading
Loading