Skip to content

Commit 03e0d4e

Browse files
committed
Add --namespace flag support to workloads
Allows specifying a target Kubernetes namespace for workload commands (create, list, delete, wait). This avoids hardcoding the 'default' namespace and enables running workloads in custom namespaces. Also updates setup_k8s_service_accounts to create service accounts and role bindings in the specified namespace.
1 parent 337e676 commit 03e0d4e

7 files changed

Lines changed: 70 additions & 50 deletions

File tree

src/xpk/commands/cluster_test.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -89,6 +89,7 @@ def construct_args(**kwargs: Any) -> Namespace:
8989
project='project',
9090
zone='us-central1-a',
9191
reservation='',
92+
namespace='',
9293
on_demand=False,
9394
tpu_type=None,
9495
device_type=None,

src/xpk/commands/workload.py

Lines changed: 11 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -393,7 +393,7 @@ def workload_create(args) -> None:
393393
k8s_api_client = None
394394
if not is_dry_run():
395395
k8s_api_client = setup_k8s_env(args)
396-
setup_k8s_service_accounts()
396+
setup_k8s_service_accounts(args.namespace)
397397

398398
workload_exists = check_if_workload_exists(args)
399399

@@ -786,7 +786,8 @@ def workload_create(args) -> None:
786786
)
787787

788788
tmp = write_tmp_file(yml_string)
789-
command = f'kubectl apply -f {str(tmp)}'
789+
ns_arg = f' -n {args.namespace}' if args.namespace else ''
790+
command = f'kubectl apply -f {str(tmp)}{ns_arg}'
790791
return_code = run_command_with_updates(command, 'Creating Workload')
791792

792793
if return_code != 0:
@@ -835,7 +836,7 @@ def workload_create(args) -> None:
835836
" python -c 'import pathwaysutils; import jax; print(jax.devices())'"
836837
)
837838
pathways_proxy_link = (
838-
f'https://console.cloud.google.com/kubernetes/job/{get_cluster_location(args.project, args.cluster, args.zone)}/{args.cluster}/default/{args.workload}-proxy-0/details?project={args.project}'
839+
f'https://console.cloud.google.com/kubernetes/job/{get_cluster_location(args.project, args.cluster, args.zone)}/{args.cluster}/{args.namespace or "default"}/{args.workload}-proxy-0/details?project={args.project}'
839840
)
840841
xpk_print(
841842
'Follow the proxy here:'
@@ -850,15 +851,18 @@ def workload_create(args) -> None:
850851
xpk_print(
851852
'Follow your workload here:'
852853
# pylint: disable=line-too-long
853-
f' https://console.cloud.google.com/kubernetes/service/{get_cluster_location(args.project, args.cluster, args.zone)}/{args.cluster}/default/{args.workload}/details?project={args.project}'
854+
f' https://console.cloud.google.com/kubernetes/service/{get_cluster_location(args.project, args.cluster, args.zone)}/{args.cluster}/{args.namespace or "default"}/{args.workload}/details?project={args.project}'
854855
)
855856
duration_of_logs = 'P1D' # Past 1 Day
857+
ns_log_filter = (
858+
f'resource.labels.namespace_name="{args.namespace or "default"}"\n'
859+
)
856860
log_filter = (
857861
'resource.type="k8s_container"\n'
858862
f'resource.labels.project_id="{args.project}"\n'
859863
f'resource.labels.location="{get_cluster_location(args.project, args.cluster, args.zone)}"\n'
860864
f'resource.labels.cluster_name="{args.cluster}"\n'
861-
'resource.labels.namespace_name="default"\n'
865+
f'{ns_log_filter}'
862866
f'resource.labels.pod_name:"{args.workload}-slice-job-0-0-"\n'
863867
'severity>=DEFAULT'
864868
)
@@ -916,7 +920,8 @@ def delete_workloads(args, workloads: list[str]) -> int:
916920
task_names = []
917921
for workload in workloads:
918922
args.workload = workload
919-
command = f'kubectl delete jobset {workload} -n default'
923+
ns_arg = f'-n {args.namespace}' if args.namespace else '-n default'
924+
command = f'kubectl delete jobset {workload} {ns_arg}'
920925
task_name = f'WorkloadDelete-{workload}'
921926
commands.append(command)
922927
task_names.append(task_name)

src/xpk/core/cluster.py

Lines changed: 22 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -428,27 +428,28 @@ def get_gpu_type_from_cluster(args) -> str:
428428
return ''
429429

430430

431-
def setup_k8s_service_accounts() -> None:
431+
def setup_k8s_service_accounts(namespace: str = 'default') -> None:
432432
"""
433433
Creates/sets up SAs and the roles for them
434434
"""
435+
namespace = namespace or 'default'
435436
default_sa = 'default'
436437

437-
create_xpk_k8s_service_account()
438+
create_xpk_k8s_service_account(namespace)
438439

439-
role_name = create_pod_reader_role()
440-
create_role_binding(default_sa, role_name)
441-
create_role_binding(XPK_SA, role_name)
440+
role_name = create_pod_reader_role(namespace)
441+
create_role_binding(default_sa, role_name, namespace)
442+
create_role_binding(XPK_SA, role_name, namespace)
442443

443444

444-
def create_xpk_k8s_service_account() -> None:
445+
def create_xpk_k8s_service_account(namespace: str = 'default') -> None:
445446
k8s_core_client = k8s_client.CoreV1Api()
446447
sa = k8s_client.V1ServiceAccount(
447448
metadata=k8s_client.V1ObjectMeta(name=XPK_SA)
448449
)
449450

450451
try:
451-
k8s_core_client.read_namespaced_service_account(XPK_SA, DEFAULT_NAMESPACE)
452+
k8s_core_client.read_namespaced_service_account(XPK_SA, namespace)
452453
xpk_print(
453454
f'Service account: {XPK_SA} already exists. Skipping its creation.'
454455
)
@@ -461,7 +462,7 @@ def create_xpk_k8s_service_account() -> None:
461462
xpk_print(f'Creating a new service account: {XPK_SA}')
462463
try:
463464
k8s_core_client.create_namespaced_service_account(
464-
DEFAULT_NAMESPACE, sa, pretty=True
465+
namespace, sa, pretty=True
465466
)
466467
xpk_print(f'Created a new service account: {XPK_SA} successfully')
467468
except ApiException as e:
@@ -474,15 +475,15 @@ def create_xpk_k8s_service_account() -> None:
474475
xpk_exit(1)
475476

476477

477-
def create_pod_reader_role() -> str:
478+
def create_pod_reader_role(namespace: str = 'default') -> str:
478479
"""
479480
Creates the 'pod-reader' Role in the default namespace.
480481
"""
481482
k8s_rbac_client = k8s_client.RbacAuthorizationV1Api()
482483
role_name = 'pod-reader'
483484

484485
try:
485-
k8s_rbac_client.read_namespaced_role(role_name, DEFAULT_NAMESPACE)
486+
k8s_rbac_client.read_namespaced_role(role_name, namespace)
486487
xpk_print(f'Role: {role_name} already exists. Skipping its creation.')
487488
return role_name
488489
except ApiException as e:
@@ -491,9 +492,7 @@ def create_pod_reader_role() -> str:
491492
xpk_exit(1)
492493

493494
role = k8s_client.V1Role(
494-
metadata=k8s_client.V1ObjectMeta(
495-
name=role_name, namespace=DEFAULT_NAMESPACE
496-
),
495+
metadata=k8s_client.V1ObjectMeta(name=role_name, namespace=namespace),
497496
rules=[
498497
k8s_client.V1PolicyRule(
499498
api_groups=[''],
@@ -508,12 +507,9 @@ def create_pod_reader_role() -> str:
508507
],
509508
)
510509

511-
xpk_print(
512-
f'Attempting to create Role: {role_name} in namespace:'
513-
f' {DEFAULT_NAMESPACE}'
514-
)
510+
xpk_print(f'Attempting to create Role: {role_name} in namespace: {namespace}')
515511
try:
516-
k8s_rbac_client.create_namespaced_role(DEFAULT_NAMESPACE, role, pretty=True)
512+
k8s_rbac_client.create_namespaced_role(namespace, role, pretty=True)
517513
xpk_print(f'Successfully created Role: {role_name}')
518514
return role_name
519515
except ApiException as e:
@@ -525,7 +521,9 @@ def create_pod_reader_role() -> str:
525521
xpk_exit(1)
526522

527523

528-
def create_role_binding(sa: str, role_name: str) -> None:
524+
def create_role_binding(
525+
sa: str, role_name: str, namespace: str = 'default'
526+
) -> None:
529527
"""
530528
Creates a RoleBinding to associate the Service Account
531529
with the Role in the default namespace.
@@ -535,9 +533,7 @@ def create_role_binding(sa: str, role_name: str) -> None:
535533
role_binding_name = f'{sa}-{role_name}-binding'
536534

537535
try:
538-
k8s_rbac_client.read_namespaced_role_binding(
539-
role_binding_name, DEFAULT_NAMESPACE
540-
)
536+
k8s_rbac_client.read_namespaced_role_binding(role_binding_name, namespace)
541537
xpk_print(
542538
f'RoleBinding: {role_binding_name} already exists. Skipping its'
543539
' creation.'
@@ -550,11 +546,11 @@ def create_role_binding(sa: str, role_name: str) -> None:
550546

551547
role_binding = k8s_client.V1RoleBinding(
552548
metadata=k8s_client.V1ObjectMeta(
553-
name=role_binding_name, namespace=DEFAULT_NAMESPACE
549+
name=role_binding_name, namespace=namespace
554550
),
555551
subjects=[
556552
k8s_client.RbacV1Subject(
557-
kind='ServiceAccount', name=sa, namespace=DEFAULT_NAMESPACE
553+
kind='ServiceAccount', name=sa, namespace=namespace
558554
)
559555
],
560556
role_ref=k8s_client.V1RoleRef(
@@ -565,11 +561,11 @@ def create_role_binding(sa: str, role_name: str) -> None:
565561
xpk_print(
566562
f'Attempting to create RoleBinding: {role_binding_name} for Service'
567563
f' Account: {sa} to Role: {role_name} in namespace:'
568-
f' {DEFAULT_NAMESPACE}'
564+
f' {namespace}'
569565
)
570566
try:
571567
k8s_rbac_client.create_namespaced_role_binding(
572-
DEFAULT_NAMESPACE, role_binding, pretty=True
568+
namespace, role_binding, pretty=True
573569
)
574570
xpk_print(f'Successfully created RoleBinding: {role_binding_name} for {sa}')
575571
except ApiException as e:

src/xpk/core/pathways.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -101,11 +101,17 @@ def check_if_pathways_job_is_installed(args) -> bool:
101101

102102
def get_pathways_unified_query_link(args) -> str:
103103
"""Get the unified query link for the pathways workload."""
104+
ns_log_filter = (
105+
f'resource.labels.namespace_name="{args.namespace}"\n'
106+
if args.namespace
107+
else ''
108+
)
104109
log_filter = (
105110
'resource.type="k8s_container"\n'
106111
f'resource.labels.project_id="{args.project}"\n'
107112
f'resource.labels.location="{get_cluster_location(args.project, args.cluster, args.zone)}"\n'
108113
f'resource.labels.cluster_name="{args.cluster}"\n'
114+
f'{ns_log_filter}'
109115
f'resource.labels.pod_name:"{args.workload}-"\n'
110116
'severity>=DEFAULT'
111117
)
@@ -143,7 +149,8 @@ def try_to_delete_pathwaysjob_first(args, workloads) -> bool:
143149
task_names = []
144150
for workload in workloads:
145151
args.workload = workload
146-
command = f'kubectl delete pathwaysjob {workload} -n default'
152+
ns_arg = f'-n {args.namespace}' if args.namespace else '-n default'
153+
command = f'kubectl delete pathwaysjob {workload} {ns_arg}'
147154
task_name = f'PathwaysWorkloadDelete-{workload}'
148155
commands.append(command)
149156
task_names.append(task_name)

src/xpk/core/workload.py

Lines changed: 22 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -240,9 +240,11 @@ def _parse_workload_item(item: dict[str, Any]) -> _WorkloadListRow:
240240
def _fetch_workloads(
241241
filter_by_status: _StatusFilter,
242242
filter_by_job: Optional[str] = None,
243+
namespace: str = '',
243244
) -> tuple[int, list[_WorkloadListRow]]:
244245
"""Fetches and parses the raw workload list from the cluster."""
245-
command = 'kubectl get workloads --ignore-not-found -o=json'
246+
ns_arg = f' -n {namespace}' if namespace else ''
247+
command = f'kubectl get workloads{ns_arg} --ignore-not-found -o=json'
246248

247249
task = f'List Jobs with filter-by-status={filter_by_status.value}'
248250
if filter_by_job:
@@ -399,7 +401,9 @@ def get_workload_list(args: argparse.Namespace) -> tuple[int, str]:
399401
filter_by_job = getattr(args, 'filter_by_job', None)
400402
filter_by_status = _get_status_filter(args.filter_by_status)
401403

402-
return_code, raw_rows = _fetch_workloads(filter_by_status, filter_by_job)
404+
return_code, raw_rows = _fetch_workloads(
405+
filter_by_status, filter_by_job, args.namespace
406+
)
403407
if return_code != 0:
404408
return return_code, ''
405409

@@ -425,7 +429,8 @@ def check_if_workload_exists(args: argparse.Namespace) -> bool:
425429

426430
s = ','.join([key + ':' + value for key, value in columns.items()])
427431

428-
command = f"kubectl get workloads -o=custom-columns='{s}'"
432+
ns_arg = f' -n {args.namespace}' if args.namespace else ''
433+
command = f"kubectl get workloads{ns_arg} -o=custom-columns='{s}'"
429434
return_code, return_msg = run_command_for_value(
430435
command, 'Check if Workload Already Exists'
431436
)
@@ -442,16 +447,20 @@ def check_if_workload_exists(args: argparse.Namespace) -> bool:
442447
return False
443448

444449

445-
def _get_jobset_status(workload_name: str) -> tuple[int, str]:
450+
def _get_jobset_status(
451+
workload_name: str, namespace: str = ''
452+
) -> tuple[int, str]:
446453
"""Retrieves the current status of a given jobset workload.
447454
448455
Args:
449456
workload_name: The name of the workload to retrieve the status for.
457+
namespace: The Kubernetes namespace to retrieve the status from.
450458
451459
Returns:
452460
A tuple containing the return code of the command (0 for success) and the status string.
453461
"""
454-
status_cmd = f'kubectl get jobset {workload_name} -o json'
462+
ns_arg = f' -n {namespace}' if namespace else ''
463+
status_cmd = f'kubectl get jobset {workload_name}{ns_arg} -o json'
455464
return_code, return_value = run_command_for_value(
456465
status_cmd, 'Get jobset status'
457466
)
@@ -497,7 +506,10 @@ def wait_for_job_completion(args: argparse.Namespace) -> int:
497506
return 1
498507

499508
# Get the full workload name
500-
get_workload_name_cmd = f'kubectl get workloads | grep jobset-{args.workload}'
509+
ns_arg = f' -n {args.namespace}' if args.namespace else ''
510+
get_workload_name_cmd = (
511+
f'kubectl get workloads{ns_arg} | grep jobset-{args.workload}'
512+
)
501513
return_code, return_value = run_command_for_value(
502514
get_workload_name_cmd, 'Get full workload name'
503515
)
@@ -512,7 +524,7 @@ def wait_for_job_completion(args: argparse.Namespace) -> int:
512524
f'{timeout_val}s' if timeout_val != -1 else 'max timeout (1 week)'
513525
)
514526
wait_cmd = (
515-
'kubectl wait --for=condition=Finished'
527+
f'kubectl wait{ns_arg} --for=condition=Finished'
516528
f' workload {full_workload_name} --timeout={timeout_val}s'
517529
)
518530
return_code, return_value = run_command_for_value(
@@ -526,7 +538,7 @@ def wait_for_job_completion(args: argparse.Namespace) -> int:
526538
f'Timed out waiting for your workload after {timeout_msg}, see your'
527539
' workload here:'
528540
# pylint: disable=line-too-long
529-
f' https://console.cloud.google.com/kubernetes/service/{get_cluster_location(args.project, args.cluster, args.zone)}/{args.cluster}/default/{args.workload}/details?project={args.project}'
541+
f' https://console.cloud.google.com/kubernetes/service/{get_cluster_location(args.project, args.cluster, args.zone)}/{args.cluster}/{args.namespace or "default"}/{args.workload}/details?project={args.project}'
530542
)
531543
return 124
532544
else:
@@ -536,9 +548,9 @@ def wait_for_job_completion(args: argparse.Namespace) -> int:
536548
xpk_print(
537549
'Finished waiting for your workload, see your workload here:'
538550
# pylint: disable=line-too-long
539-
f' https://console.cloud.google.com/kubernetes/service/{get_cluster_location(args.project, args.cluster, args.zone)}/{args.cluster}/default/{args.workload}/details?project={args.project}'
551+
f' https://console.cloud.google.com/kubernetes/service/{get_cluster_location(args.project, args.cluster, args.zone)}/{args.cluster}/{args.namespace or "default"}/{args.workload}/details?project={args.project}'
540552
)
541-
return_code, return_value = _get_jobset_status(args.workload)
553+
return_code, return_value = _get_jobset_status(args.workload, args.namespace)
542554
if return_code != 0:
543555
return return_code
544556

src/xpk/parser/common.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -108,6 +108,12 @@ def add_shared_arguments(
108108
),
109109
required=required,
110110
)
111+
custom_parser_or_group.add_argument(
112+
'--namespace',
113+
type=str,
114+
default='',
115+
help='Kubernetes namespace to use. Defaults to active namespace.',
116+
)
111117
custom_parser_or_group.add_argument(
112118
'--dry-run',
113119
type=bool,

src/xpk/parser/info.py

Lines changed: 0 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -36,13 +36,6 @@ def set_info_parser(info_parser: argparse.ArgumentParser) -> None:
3636
required=True,
3737
)
3838

39-
info_optional_arguments.add_argument(
40-
'--namespace',
41-
type=str,
42-
default='',
43-
help='Namespace to which resources and queues belong',
44-
)
45-
4639
queues_flitering_group = (
4740
info_optional_arguments.add_mutually_exclusive_group()
4841
)

0 commit comments

Comments
 (0)