google-colab-cli 0.5.4__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,508 @@
1
+ # Copyright 2026 Google LLC
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import os
16
+ import subprocess
17
+ import sys
18
+ import time
19
+ import uuid
20
+ from typing import Any, Dict, Optional
21
+ import typer
22
+ from typing_extensions import Annotated
23
+
24
+ from colab_cli.client import (
25
+ Accelerator,
26
+ ColabRequestError,
27
+ PostAssignmentResponse,
28
+ Variant,
29
+ )
30
+ from colab_cli.commands.automation import INTERACTIVE_AUTOMATION_TIMEOUT_SEC
31
+ from colab_cli.utils import get_status_code
32
+ from colab_cli.state import SessionState
33
+ from colab_cli.runtime import ColabRuntime
34
+
35
+
36
+ def _is_scope_error(e: Exception) -> bool:
37
+ """True if a ColabRequestError's response body indicates a missing OAuth scope.
38
+
39
+ The frontend returns a `google.rpc.Status` with `code=7` (PERMISSION_DENIED)
40
+ and a `DebugInfo` payload mentioning `SCOPE_NOT_PERMITTED` /
41
+ "insufficient authentication scopes". Match on either substring so we
42
+ don't depend on the exact wording of one of them.
43
+ """
44
+ body = getattr(e, "response_body", None) or ""
45
+ body_str = str(body)
46
+ return (
47
+ "SCOPE_NOT_PERMITTED" in body_str
48
+ or "insufficient authentication scopes" in body_str
49
+ )
50
+
51
+
52
+ def _scope_remediation_message(provider) -> str:
53
+ """User-facing remediation hint, tailored per auth provider."""
54
+ # Importing locally to avoid a circular import at module load time.
55
+ from colab_cli.auth import AuthProvider
56
+
57
+ common = (
58
+ "The Colab keep-alive RPC requires the "
59
+ "'https://www.googleapis.com/auth/colaboratory' OAuth scope."
60
+ )
61
+ if provider == AuthProvider.ADC:
62
+ return (
63
+ f"{common}\n"
64
+ "Re-authenticate ADC with both userinfo.email (required by the "
65
+ "Colab session backend at colab.research.google.com) and "
66
+ "colaboratory (required by the runtime service at "
67
+ "colab.pa.googleapis.com). The cloud-platform and openid scopes "
68
+ "are required by gcloud itself:\n"
69
+ " gcloud auth application-default login \\\n"
70
+ " --scopes=openid,"
71
+ "https://www.googleapis.com/auth/cloud-platform,"
72
+ "https://www.googleapis.com/auth/userinfo.email,"
73
+ "https://www.googleapis.com/auth/colaboratory\n"
74
+ "Then re-run `colab new`."
75
+ )
76
+ # OAuth2 (and any future provider) fallback.
77
+ return (
78
+ f"{common}\n"
79
+ "Delete the cached token at ~/.config/colab-cli/token.json and "
80
+ "re-run `colab new` to trigger a fresh consent flow that includes "
81
+ "the colaboratory scope."
82
+ )
83
+
84
+
85
+ def _hardware_label(accelerator: str) -> str:
86
+ """`NONE` -> `CPU`; everything else passes through."""
87
+ return "CPU" if accelerator == "NONE" else accelerator
88
+
89
+
90
+ def _format_session_line(
91
+ name: str,
92
+ endpoint: str,
93
+ accelerator: str,
94
+ variant: str,
95
+ status: Optional[str] = None,
96
+ ) -> str:
97
+ """Single source of truth for session display lines.
98
+
99
+ Format: ``[name] endpoint | Hardware: X | Variant: Y[ | Status: Z]``.
100
+ Use ``"?"`` as the name for orphaned server-side assignments with no local
101
+ state.
102
+ """
103
+ parts = [
104
+ f"[{name}] {endpoint}",
105
+ f"Hardware: {_hardware_label(accelerator)}",
106
+ f"Variant: {variant}",
107
+ ]
108
+ if status is not None:
109
+ parts.append(f"Status: {status}")
110
+ return " | ".join(parts)
111
+
112
+
113
+ def new(
114
+ session: Annotated[
115
+ Optional[str], typer.Option("-s", "--session", help="Session name")
116
+ ] = None,
117
+ tpu: Annotated[
118
+ Optional[str],
119
+ typer.Option(
120
+ help="TPU accelerator variant. Supported: v5e1, v6e1.",
121
+ ),
122
+ ] = None,
123
+ gpu: Annotated[
124
+ Optional[str],
125
+ typer.Option(
126
+ help=(
127
+ "GPU accelerator variant. Supported: T4, L4, G4, H100, A100."
128
+ "\n\nIf omitted (along with --tpu), a CPU runtime is created."
129
+ "\n\nAvailability varies by Colab subscription tier."
130
+ ),
131
+ ),
132
+ ] = None,
133
+ ):
134
+ """Create a new session"""
135
+ from colab_cli.common import state
136
+
137
+ name = session or uuid.uuid4().hex[:6]
138
+ variant = Variant.DEFAULT
139
+ accelerator = Accelerator.NONE
140
+
141
+ if tpu:
142
+ variant = Variant.TPU
143
+ accelerator = Accelerator.V5E1 if tpu.lower() == "v5e1" else Accelerator.V6E1
144
+ elif gpu:
145
+ variant = Variant.GPU
146
+ mapping = {
147
+ "a100": Accelerator.A100,
148
+ "h100": Accelerator.H100,
149
+ "l4": Accelerator.L4,
150
+ "t4": Accelerator.T4,
151
+ "g4": Accelerator.G4,
152
+ }
153
+ accelerator = mapping.get(gpu.lower(), Accelerator.A100)
154
+
155
+ typer.echo(f"[colab] Creating session '{name}'...")
156
+ try:
157
+ res = state.client.assign(
158
+ uuid.uuid4(), variant=variant, accelerator=accelerator
159
+ )
160
+ except ColabRequestError as e:
161
+ # The Colab backend returns 400 when the caller is not entitled to the
162
+ # requested accelerator (e.g. no A100 quota). Translate that to a
163
+ # friendly, actionable message instead of a raw traceback. We only
164
+ # interpret it this way when an accelerator was actually requested;
165
+ # otherwise we re-raise so the user sees the real cause.
166
+ if get_status_code(e) == 400 and accelerator != Accelerator.NONE:
167
+ typer.echo(
168
+ f"[colab] Backend rejected accelerator '{accelerator.value}'. "
169
+ "You may not have quota or entitlement for this accelerator on "
170
+ "your account. Try a different one (e.g. --gpu T4) or omit "
171
+ "--gpu/--tpu for a CPU runtime.",
172
+ err=True,
173
+ )
174
+ raise typer.Exit(code=1)
175
+ raise
176
+
177
+ if isinstance(res, PostAssignmentResponse):
178
+ token = res.runtime_proxy_info.token
179
+ url = res.runtime_proxy_info.url
180
+ endpoint = res.endpoint
181
+ else:
182
+ token = (
183
+ res.runtime_proxy_info.token
184
+ if hasattr(res, "runtime_proxy_info")
185
+ else getattr(res, "runtime_proxy_token", "")
186
+ )
187
+ url = res.runtime_proxy_info.url if hasattr(res, "runtime_proxy_info") else ""
188
+ endpoint = res.endpoint
189
+
190
+ # Importing locally to avoid a top-level circular import via auth.
191
+
192
+ s = SessionState(
193
+ name=name,
194
+ token=token,
195
+ url=url,
196
+ endpoint=endpoint,
197
+ variant=variant.value,
198
+ accelerator=accelerator.value,
199
+ )
200
+
201
+ # Pre-flight the keep-alive RPC once. If it returns 403 SCOPE_NOT_PERMITTED
202
+ # we know the daemon will fail and the VM would be idle-pruned. Catch
203
+ # it now so we (a) never leak a billable assignment, (b) surface an
204
+ # actionable remediation instead of a "session quietly disappeared".
205
+ try:
206
+ state.client.keep_alive_assignment(endpoint)
207
+ except ColabRequestError as e:
208
+ if get_status_code(e) == 403 and _is_scope_error(e):
209
+ typer.echo(
210
+ "[colab] Keep-alive pre-flight failed: your OAuth "
211
+ "credentials are missing the 'colaboratory' scope, which "
212
+ "is required by the Colab RuntimeService.\n",
213
+ err=True,
214
+ )
215
+ typer.echo(_scope_remediation_message(state.auth_provider), err=True)
216
+ # Don't leak the assignment we just created.
217
+ try:
218
+ state.client.unassign(endpoint)
219
+ except Exception:
220
+ pass
221
+ raise typer.Exit(code=1)
222
+ # Other failures: don't block session creation — the daemon will
223
+ # retry and log via the existing keep_alive_error event path.
224
+
225
+ # Persist the session BEFORE spawning the daemon so the daemon's
226
+ # initial `state.store.get(session_name)` check doesn't race and
227
+ # exit with `reason=session_not_found`. We re-persist below to also
228
+ # capture the daemon PID.
229
+ state.store.add(s)
230
+ s.keep_alive_pid = spawn_keep_alive(
231
+ endpoint,
232
+ name,
233
+ auth_provider=state.auth_provider,
234
+ config_path=state.config_path,
235
+ )
236
+
237
+ state.store.add(s)
238
+ state.history.log_event(
239
+ name,
240
+ "session_created",
241
+ {
242
+ "endpoint": endpoint,
243
+ "variant": variant.value,
244
+ "accelerator": accelerator.value,
245
+ },
246
+ )
247
+ typer.echo("[colab] Session READY.")
248
+
249
+
250
+ def restart(
251
+ session: Annotated[
252
+ Optional[str], typer.Option("-s", "--session", help="Session name")
253
+ ] = None,
254
+ ):
255
+ from colab_cli.common import state
256
+
257
+ name = state.resolve_session(session)
258
+ s = state.store.get(name)
259
+
260
+ def on_started(kid):
261
+ s.kernel_id = kid
262
+ state.store.add(s)
263
+
264
+ def on_sess_started(sid):
265
+ s.session_id = sid
266
+ state.store.add(s)
267
+
268
+ runtime = ColabRuntime(
269
+ s.url,
270
+ s.token,
271
+ kernel_id=s.kernel_id,
272
+ session_id=s.session_id,
273
+ on_kernel_started=on_started,
274
+ on_session_started=on_sess_started,
275
+ )
276
+
277
+ runtime.restart(timeout=INTERACTIVE_AUTOMATION_TIMEOUT_SEC)
278
+
279
+
280
+ def sessions_command():
281
+ """List all active sessions"""
282
+ from colab_cli.common import state
283
+
284
+ sessions, assignments = state.sync_sessions()
285
+ if not assignments:
286
+ typer.echo("[colab] No active sessions found on server.")
287
+ return
288
+
289
+ # Build endpoint -> local-name lookup so we can lead with the friendly name.
290
+ name_by_endpoint = {s.endpoint: s.name for s in sessions.values()}
291
+ for a in assignments:
292
+ name = name_by_endpoint.get(a.endpoint, "?")
293
+ # `a.variant` is an int-valued AssignmentVariant (DEFAULT=0/GPU=1/TPU=2);
294
+ # its `.name` matches the user-facing string Variant enum, which is what
295
+ # `status` shows for locally-tracked sessions.
296
+ typer.echo(
297
+ _format_session_line(
298
+ name=name,
299
+ endpoint=a.endpoint,
300
+ accelerator=a.accelerator.value,
301
+ variant=a.variant.name,
302
+ )
303
+ )
304
+
305
+
306
+ def _print_status_for(s: SessionState) -> None:
307
+ """Print one session's status line plus optional last-execution detail."""
308
+ status = f"BUSY ({s.running})" if s.running else "IDLE"
309
+ typer.echo(
310
+ _format_session_line(
311
+ name=s.name,
312
+ endpoint=s.endpoint,
313
+ accelerator=s.accelerator,
314
+ variant=s.variant,
315
+ status=status,
316
+ )
317
+ )
318
+ if s.last_execution:
319
+ exec_file, exec_cell, exec_time = s.last_execution
320
+ cell_str = f" | Cell: {exec_cell}" if exec_cell else ""
321
+ typer.echo(f" Last Execution: {exec_file}{cell_str} at {exec_time}")
322
+
323
+
324
+ def status(
325
+ session: Annotated[
326
+ Optional[str], typer.Option("-s", "--session", help="Session name")
327
+ ] = None,
328
+ ):
329
+ """Show session status"""
330
+ from colab_cli.common import state
331
+
332
+ local_sessions, _ = state.sync_sessions()
333
+ if session:
334
+ s = state.store.get(session)
335
+ if s:
336
+ _print_status_for(s)
337
+ else:
338
+ typer.echo(f"[colab] Session '{session}' not found.")
339
+ return
340
+
341
+ if not local_sessions:
342
+ typer.echo("[colab] No active sessions.")
343
+ return
344
+ for s in local_sessions.values():
345
+ _print_status_for(s)
346
+
347
+
348
+ def stop(
349
+ session: Annotated[
350
+ Optional[str], typer.Option("-s", "--session", help="Session name")
351
+ ] = None,
352
+ ):
353
+ """Stop a session"""
354
+ from colab_cli.common import state
355
+
356
+ name = state.resolve_session(session)
357
+ s = state.store.get(name)
358
+ if not s:
359
+ typer.echo(f"[colab] Session '{name}' not found.")
360
+ return
361
+
362
+ typer.echo(f"[colab] Stopping session '{name}'...")
363
+ if s.keep_alive_pid:
364
+ from colab_cli.common import kill_process
365
+
366
+ kill_process(s.keep_alive_pid)
367
+
368
+ try:
369
+ runtime = ColabRuntime(s.url, s.token, kernel_id=s.kernel_id)
370
+ runtime.stop(shutdown_kernel=True)
371
+ except Exception:
372
+ pass
373
+
374
+ state.client.unassign(s.endpoint)
375
+ state.store.remove(name)
376
+ state.history.log_event(name, "session_terminated", {"reason": "user_requested"})
377
+ typer.echo("[colab] Session terminated.")
378
+
379
+
380
+ def spawn_keep_alive(
381
+ endpoint: str, session_name: str, auth_provider=None, config_path=None
382
+ ):
383
+ """Spawns a detached keep-alive process.
384
+
385
+ Both `auth_provider` and `config_path` are propagated as global flags
386
+ so the detached child uses the same authentication strategy AND the
387
+ same session state file as the parent that invoked `colab new`.
388
+ Without this, the child inherits Typer's defaults (`--auth=oauth2`,
389
+ `--config=~/.config/colab-cli/sessions.json`), which causes:
390
+ (a) wrong auth backend, and
391
+ (b) the daemon's `state.store.get(session_name)` check finds nothing
392
+ and exits with `reason=session_not_found` when the parent used
393
+ `--config` to write to a non-default path.
394
+ """
395
+ cmd = [sys.executable, "-m", "colab_cli.cli"]
396
+ if auth_provider is not None:
397
+ cmd.append(f"--auth={auth_provider.value}")
398
+ if config_path is not None:
399
+ cmd.extend(["--config", config_path])
400
+ cmd.extend(["keep-alive", endpoint, session_name])
401
+ # Detach process
402
+ kwargs = {}
403
+ if sys.platform != "win32":
404
+ kwargs["start_new_session"] = True
405
+ else:
406
+ # https://stackoverflow.com/questions/1356540/how-can-i-make-a-python-script-run-in-the-background-as-a-service-on-windows
407
+ CREATE_NEW_PROCESS_GROUP = 0x00000200
408
+ DETACHED_PROCESS = 0x00000008
409
+ kwargs["creationflags"] = DETACHED_PROCESS | CREATE_NEW_PROCESS_GROUP
410
+
411
+ p = subprocess.Popen(
412
+ cmd,
413
+ stdout=subprocess.DEVNULL,
414
+ stderr=subprocess.DEVNULL,
415
+ stdin=subprocess.DEVNULL,
416
+ **kwargs,
417
+ )
418
+ return p.pid
419
+
420
+
421
+ def keep_alive(
422
+ endpoint: Annotated[str, typer.Argument(help="Endpoint ID")],
423
+ session_name: Annotated[str, typer.Argument(help="Session name")],
424
+ ):
425
+ """Hidden command to run keep-alive loop. Terminate after 24h."""
426
+ from colab_cli.common import state
427
+
428
+ state.history.log_event(
429
+ session_name,
430
+ "keep_alive_started",
431
+ {"endpoint": endpoint, "pid": os.getpid()},
432
+ )
433
+
434
+ start_time = time.time()
435
+ # 24 hours limit
436
+ max_duration = 24 * 3600
437
+ consecutive_4xx = 0
438
+ iterations = 0
439
+ last_error: Optional[Dict[str, Any]] = None
440
+
441
+ reason = "time_limit_reached"
442
+ extra: Dict[str, Any] = {}
443
+ while time.time() - start_time < max_duration:
444
+ iterations += 1
445
+ # Check if session still exists in local state
446
+ s = state.store.get(session_name)
447
+ if not s:
448
+ reason = "session_not_found"
449
+ break
450
+ if s.endpoint != endpoint:
451
+ reason = "endpoint_mismatch"
452
+ extra["expected_endpoint"] = endpoint
453
+ extra["actual_endpoint"] = s.endpoint
454
+ break
455
+
456
+ try:
457
+ state.client.keep_alive_assignment(endpoint)
458
+ consecutive_4xx = 0
459
+ last_error = None
460
+ except Exception as e:
461
+ code = get_status_code(e)
462
+ response_body = getattr(e, "response_body", None)
463
+ err_info = {
464
+ "status_code": code,
465
+ "error_type": type(e).__name__,
466
+ "error": str(e)[:500],
467
+ "response_body": (str(response_body)[:1000] if response_body else None),
468
+ }
469
+ last_error = err_info
470
+ state.history.log_event(
471
+ session_name,
472
+ "keep_alive_error",
473
+ {
474
+ **err_info,
475
+ "iteration": iterations,
476
+ "consecutive_4xx": consecutive_4xx
477
+ + (1 if code is not None and 400 <= code < 500 else 0),
478
+ },
479
+ )
480
+ if code is not None and 400 <= code < 500:
481
+ consecutive_4xx += 1
482
+ if consecutive_4xx >= 2:
483
+ reason = "consecutive_4xx_errors"
484
+ break
485
+ else:
486
+ # For other errors (network), we retry and don't count as 4xx
487
+ pass
488
+
489
+ time.sleep(60)
490
+
491
+ payload: Dict[str, Any] = {
492
+ "reason": reason,
493
+ "iterations": iterations,
494
+ "duration_seconds": round(time.time() - start_time, 2),
495
+ }
496
+ if last_error is not None:
497
+ payload["last_error"] = last_error
498
+ payload.update(extra)
499
+ state.history.log_event(session_name, "keep_alive_stopped", payload)
500
+
501
+
502
+ def register(app: typer.Typer):
503
+ app.command()(new)
504
+ app.command(name="sessions")(sessions_command)
505
+ app.command()(restart)
506
+ app.command()(status)
507
+ app.command()(stop)
508
+ app.command(hidden=True)(keep_alive)