44from typing import Optional
55from absl import logging
66import re
7+ from enum import Flag , auto
78
89from airflow .decorators import task
910from airflow .exceptions import AirflowFailException
1011from google .cloud import logging as logging_api
1112from xlml .apis import gcs
1213
14+ class FilterMode (Flag ):
15+ textPayload = auto ()
16+ jsonPayload_message = auto ()
17+
1318
1419@task
1520def generate_timestamp ():
@@ -60,11 +65,8 @@ def validate_checkpoint_at_steps_are_saved(
6065 location = location ,
6166 cluster_name = cluster_name ,
6267 pod_pattern = pod_pattern ,
63- text_filter = (
64- "\" 'event_type': 'save'\" AND "
65- f'(textPayload=~"{ log_pattern } " OR '
66- f'jsonPayload.message=~"{ log_pattern } ")'
67- ),
68+ text_filters = [r"'event_type': 'save'" ],
69+ filter_mode = FilterMode .textPayload | FilterMode .jsonPayload_message ,
6870 start_time = start_time ,
6971 end_time = end_time ,
7072 )
@@ -137,7 +139,8 @@ def validate_log_with_gcs(
137139 namespace : str = "default" ,
138140 pod_pattern : str = ".*" ,
139141 container_name : Optional [str ] = None ,
140- text_filter : Optional [str ] = None ,
142+ text_filters : Optional [list [str ]] = None ,
143+ filter_mode : Optional ["FilterMode" ] = None ,
141144 start_time : Optional [datetime ] = None ,
142145 end_time : Optional [datetime ] = None ,
143146) -> None :
@@ -158,8 +161,8 @@ def validate_log_with_gcs(
158161 namespace: The Kubernetes namespace. Defaults to "default".
159162 pod_pattern: A glob pattern to match pod names. Defaults to "*".
160163 container_name: An optional container name to filter logs by.
161- text_filter : An optional string to filter log entries by their
162- `textPayload` .
164+ text_filters : An optional list of strings to filter log entries.
165+ filter_mode: A FilterMode flag to specify which payload fields to search .
163166 start_time: The start time for log retrieval.
164167 end_time: The end time for log retrieval.
165168
@@ -181,7 +184,8 @@ def validate_log_with_gcs(
181184 namespace = namespace ,
182185 pod_pattern = pod_pattern ,
183186 container_name = container_name ,
184- text_filter = f'textPayload=~"{ text_filter } "' ,
187+ text_filters = text_filters ,
188+ filter_mode = filter_mode ,
185189 start_time = start_time ,
186190 end_time = end_time ,
187191 )
@@ -352,7 +356,8 @@ def list_log_entries(
352356 namespace : str = "default" ,
353357 pod_pattern : str = ".*" ,
354358 container_name : Optional [str ] = None ,
355- text_filter : Optional [str ] = None ,
359+ text_filters : Optional [list [str ]] = None ,
360+ filter_mode : Optional ["FilterMode" ] = None ,
356361 start_time : Optional [datetime ] = None ,
357362 end_time : Optional [datetime ] = None ,
358363) -> list [logging_api .LogEntry ]:
@@ -373,8 +378,8 @@ def list_log_entries(
373378 namespace: Kubernetes namespace (defaults to "default")
374379 pod_pattern: Pattern to match pod names (defaults to "*")
375380 container_name: Optional container name to filter logs
376- text_filter : Optional comma-separated string to
377- filter log entries by textPayload content
381+ text_filters : Optional list of strings to filter log entries
382+ filter_mode: Optional FilterMode to specify which payload fields to search
378383 start_time: Optional start time for log retrieval
379384 (defaults to 12 hours ago)
380385 end_time: Optional end time for log retrieval (defaults to now)
@@ -408,8 +413,15 @@ def list_log_entries(
408413
409414 if container_name :
410415 conditions .append (f'resource.labels.container_name="{ container_name } "' )
411- if text_filter :
412- conditions .append (f"{ text_filter } " )
416+ if text_filters and filter_mode is not None :
417+ filters = []
418+ for txt in text_filters :
419+ if FilterMode .textPayload in filter_mode :
420+ filters .append (f'textPayload=~"{ txt } "' )
421+ if FilterMode .jsonPayload_message in filter_mode :
422+ filters .append (f'jsonPayload.message=~"{ txt } "' )
423+ if filters :
424+ conditions .append (f"({ ' OR ' .join (filters )} )" )
413425
414426 log_filter = " AND " .join (conditions )
415427
@@ -425,7 +437,8 @@ def validate_log_exist(
425437 namespace : str = "default" ,
426438 pod_pattern : str = ".*" ,
427439 container_name : Optional [str ] = None ,
428- text_filter : Optional [str ] = None ,
440+ text_filters : Optional [list [str ]] = None ,
441+ filter_mode : Optional ["FilterMode" ] = None ,
429442 start_time : Optional [datetime ] = None ,
430443 end_time : Optional [datetime ] = None ,
431444) -> None :
@@ -438,7 +451,8 @@ def validate_log_exist(
438451 namespace = namespace ,
439452 pod_pattern = pod_pattern ,
440453 container_name = container_name ,
441- text_filter = text_filter ,
454+ text_filters = text_filters ,
455+ filter_mode = filter_mode ,
442456 start_time = start_time ,
443457 end_time = end_time ,
444458 )
@@ -460,21 +474,24 @@ def validate_restored_correct_checkpoint(
460474) -> None :
461475 """Validate the restored step is in the expected range."""
462476
477+ reg_save_event = r"'event_type': 'save'"
478+ reg_restor_event = r"'event_type': '(emergency_)?restore'"
479+ reg_restoring = r"restoring from this run's directory step (\d+)"
480+
463481 entries = list_log_entries (
464482 project_id = project_id ,
465483 location = location ,
466484 cluster_name = cluster_name ,
467485 namespace = "default" ,
468486 pod_pattern = pod_pattern ,
469- text_filter = (
470- "(textPayload:\" 'event_type'\" OR jsonPayload.message:\" 'event_type'\" )"
471- ),
487+ text_filters = [reg_save_event , reg_restor_event , reg_restoring ],
488+ filter_mode = FilterMode .textPayload | FilterMode .jsonPayload_message ,
472489 start_time = start_time ,
473490 end_time = end_time ,
474491 )
475492
476493 if not entries :
477- raise AirflowFailException ("No event_type found in the log." )
494+ raise AirflowFailException ("No event_type or restore log found in the log." )
478495
479496 local_saved_steps_before_restore = []
480497 for entry in entries :
@@ -488,7 +505,7 @@ def validate_restored_correct_checkpoint(
488505 logging .warning (f"Could not extract message from log entry: { entry } " )
489506 continue
490507
491- if re .search (r"'event_type': 'save'" , message ):
508+ if re .search (reg_save_event , message ):
492509 saved_step_match = re .search (r"'step': (\d+)" , message )
493510 if not saved_step_match :
494511 raise AirflowFailException (
@@ -497,15 +514,15 @@ def validate_restored_correct_checkpoint(
497514
498515 local_saved_steps_before_restore .append (int (saved_step_match .group (1 )))
499516
500- elif re .search (r"'event_type': '(emergency_)?restore'" , message ):
517+ elif re .search (reg_restor_event , message ) or re . search ( reg_restoring , message ):
501518 logging .info ("Found restore event: %s" , message )
502519 logging .info (
503520 "Saved steps before restore: %s" , local_saved_steps_before_restore
504521 )
505522
506523 restored_step_match = re .search (
507524 r"'step':\s*(?:np\.int32\()?(\d+)" , message
508- )
525+ ) or re . search ( reg_restoring , message )
509526 restored_step = (
510527 int (restored_step_match .group (1 )) if restored_step_match else None
511528 )
@@ -585,7 +602,8 @@ def validate_replicator_gcs_restore_log(
585602 namespace = namespace ,
586603 pod_pattern = pod_pattern ,
587604 container_name = container_name ,
588- text_filter = "Restoring from backup" ,
605+ text_filters = ["Restoring from backup" ],
606+ filter_mode = FilterMode .textPayload | FilterMode .jsonPayload_message ,
589607 start_time = start_time ,
590608 end_time = end_time ,
591609 )
@@ -680,7 +698,8 @@ def validate_replicator_gcs_backup_log(
680698 namespace = namespace ,
681699 pod_pattern = pod_pattern ,
682700 container_name = container_name ,
683- text_filter = f'textPayload=~"{ step_regex } "' ,
701+ text_filters = [step_regex ],
702+ filter_mode = FilterMode .textPayload ,
684703 start_time = start_time ,
685704 end_time = end_time ,
686705 )
@@ -737,7 +756,8 @@ def validate_checkpoints_save_regular_axlearn(
737756 location = location ,
738757 cluster_name = cluster_name ,
739758 pod_pattern = pod_pattern ,
740- text_filter = f'jsonPayload.message=~"{ log_pattern } "' ,
759+ text_filters = [log_pattern ],
760+ filter_mode = FilterMode .jsonPayload_message ,
741761 start_time = start_time ,
742762 end_time = end_time ,
743763 )
0 commit comments