import json
import os
from typing import Optional
import gradio as gr
import numpy as np
import pandas as pd
from data_loader import (
ARTIFACT_LABEL_MAP,
get_annotations,
get_annotations_in_window,
get_channel_names,
get_recording_display_list,
get_recording_info,
get_store_metadata,
get_tuar_recordings,
preload_all_annotations,
read_signal_window,
reset_s3fs,
)
from visualizer import (
build_annotation_summary,
build_artifact_legend,
build_eeg_figure,
build_metadata_html,
)
WINDOW_PADDING_S = 2.0
MAX_WINDOW_S = 15.0
MIN_WINDOW_S = 2.0
ALL_ARTIFACT_TYPES = [
"Eye Movement", "Muscle", "Electrode Pop",
"Chewing", "Shiver", "Artifact (Generic)", "Background",
"Eye Movement + Muscle", "Muscle + Electrode Pop",
"Eye Movement + Electrode Pop", "Eye Movement + Chewing",
"Chewing + Electrode Pop", "Chewing + Muscle",
"Eye Movement + Shiver", "Shiver + Electrode Pop",
]
recordings_df: pd.DataFrame = pd.DataFrame()
# Maps recording_label -> list of annotations
annotations_index: dict[str, list[dict]] = {}
# Maps recording_label -> info dict
info_index: dict[str, dict] = {}
# Maps recording_label -> store metadata
meta_index: dict[str, dict] = {}
def check_aws_credentials() -> bool:
return bool(os.environ.get("AWS_ACCESS_KEY_ID") and os.environ.get("AWS_SECRET_ACCESS_KEY"))
def save_credentials(access_key: str, secret_key: str, region: str) -> str:
if not access_key.strip() or not secret_key.strip():
return '
Both Access Key and Secret Key are required.
'
os.environ["AWS_ACCESS_KEY_ID"] = access_key.strip()
os.environ["AWS_SECRET_ACCESS_KEY"] = secret_key.strip()
os.environ["AWS_DEFAULT_REGION"] = region.strip() or "us-east-1"
reset_s3fs()
env_path = os.path.join(os.path.dirname(__file__), ".env")
with open(env_path, "w") as f:
f.write(f"AWS_ACCESS_KEY_ID={access_key.strip()}\n")
f.write(f"AWS_SECRET_ACCESS_KEY={secret_key.strip()}\n")
f.write(f"AWS_DEFAULT_REGION={region.strip() or 'us-east-1'}\n")
return 'Credentials saved.
'
_annotations_loaded = False
def init_recordings():
global recordings_df
try:
recordings_df = get_tuar_recordings()
if len(recordings_df) == 0:
return 'No TUAR recordings found.
'
return (
f'Loaded {len(recordings_df)} TUAR recordings. '
f'Select an artifact type to begin.
'
)
except Exception as e:
return f'Error loading manifest: {e}
'
def _ensure_annotations_loaded(progress=None):
global _annotations_loaded
if _annotations_loaded:
return
if progress:
progress(0.1, desc="Fetching artifact annotations from S3...")
preload_all_annotations(recordings_df)
if progress:
progress(0.7, desc="Building index...")
for _, row in recordings_df.iterrows():
info = get_recording_info(row)
canonical_uri = info.get("canonical_uri", "")
if not canonical_uri:
continue
anns = get_annotations(canonical_uri, source_uri=info.get("archival_uri", ""))
if not anns:
continue
rec_key = info["recording_id"]
annotations_index[rec_key] = anns
info_index[rec_key] = info
_annotations_loaded = True
if progress:
progress(1.0, desc="Done!")
def on_artifact_type_selected(artifact_type: str, progress=gr.Progress()):
"""Fetch annotations on first use, then filter by type."""
if not artifact_type or recordings_df.empty:
return (
gr.Dropdown(choices=[], value=None),
'No recordings loaded.
',
)
_ensure_annotations_loaded(progress)
if not annotations_index:
return (
gr.Dropdown(choices=[], value=None),
'No annotations found. Check AWS credentials.
',
)
matching = []
for rec_key, anns in annotations_index.items():
type_anns = [a for a in anns if a["label"] == artifact_type]
if type_anns:
info = info_index[rec_key]
label = (
f"{rec_key[:8]}… | "
f"subj={info.get('subject','?')} | "
f"ses={info.get('session','?')} | "
f"{len(type_anns)} instance(s) | "
f"dur={info.get('duration_s',0):.0f}s"
)
matching.append((label, rec_key))
if not matching:
return (
gr.Dropdown(choices=[], value=None),
f'No recordings contain {artifact_type}.
',
)
choices = [m[0] for m in matching]
return (
gr.Dropdown(
choices=choices, value=choices[0],
label=f"Recordings with {artifact_type} ({len(choices)} found)",
),
f'{len(choices)} recordings with {artifact_type}.
',
)
def _find_rec_key(recording_label: str) -> Optional[str]:
prefix = recording_label.split("…")[0] if "…" in recording_label else recording_label[:8]
for key in annotations_index:
if key.startswith(prefix):
return key
return None
def on_recording_selected(artifact_type: str, recording_label: str):
"""Show artifact instances for the selected recording + type."""
if not recording_label or not artifact_type:
return (
gr.Dropdown(choices=[], value=None),
'',
None,
'',
gr.CheckboxGroup(choices=[], value=[]),
)
rec_key = _find_rec_key(recording_label)
if not rec_key:
return (
gr.Dropdown(choices=[], value=None),
'Recording not found in index.
',
None,
'',
gr.CheckboxGroup(choices=[], value=[]),
)
anns = annotations_index.get(rec_key, [])
info = info_index.get(rec_key, {})
type_anns = [a for a in anns if a["label"] == artifact_type]
if not type_anns:
return (
gr.Dropdown(choices=[], value=None),
build_metadata_html(info),
None,
'No instances found.
',
gr.CheckboxGroup(choices=[], value=[]),
)
canonical_uri = info.get("canonical_uri", "")
channel_names = info.get("channel_names", [])
if isinstance(channel_names, str):
import json as _json
try:
channel_names = _json.loads(channel_names)
except Exception:
channel_names = []
if not isinstance(channel_names, list):
channel_names = list(channel_names)
store_meta = {
"channel_names": channel_names,
"sampling_rate_hz": info.get("sampling_rate_hz", 250),
"duration_s": info.get("duration_s", 0),
}
meta_index[rec_key] = store_meta
choices = []
for i, inst in enumerate(type_anns):
ch = inst.get("channel", "all")
choices.append(
f"#{i+1} | {inst['onset_s']:.1f}s – {inst['end_s']:.1f}s | "
f"dur={inst['duration_s']:.1f}s | ch={ch}"
)
all_channels = store_meta.get("channel_names", [])
return (
gr.Dropdown(choices=choices, value=choices[0],
label=f"{artifact_type} instances ({len(choices)})"),
build_metadata_html(info),
None,
build_annotation_summary(type_anns),
gr.CheckboxGroup(choices=all_channels, value=[], label=f"Channels ({len(all_channels)})"),
)
def on_instance_selected(artifact_type: str, recording_label: str, instance_label: str, selected_channels: list[str]):
"""Render the EEG plot for the selected artifact instance."""
if not instance_label or not recording_label:
return (
build_eeg_figure(np.zeros((1, 100)), ["Pick an instance"], 256.0, title="Select an artifact instance"),
'',
gr.CheckboxGroup(),
)
try:
idx = int(instance_label.split("|")[0].strip().replace("#", "")) - 1
except (ValueError, IndexError):
return (
build_eeg_figure(np.zeros((1, 100)), ["Error"], 256.0, title="Parse error"),
'',
gr.CheckboxGroup(),
)
rec_key = _find_rec_key(recording_label) if recording_label else None
anns = annotations_index.get(rec_key, []) if rec_key else []
info = info_index.get(rec_key, {}) if rec_key else {}
store_meta = meta_index.get(rec_key, {}) if rec_key else {}
type_anns = [a for a in anns if a["label"] == artifact_type]
if idx < 0 or idx >= len(type_anns):
return (
build_eeg_figure(np.zeros((1, 100)), ["Error"], 256.0, title="Instance not found"),
'',
gr.CheckboxGroup(),
)
artifact = type_anns[idx]
canonical_uri = info.get("canonical_uri", "")
all_channels = store_meta.get("channel_names", [])
sfreq = store_meta.get("sampling_rate_hz", 250.0)
duration = store_meta.get("duration_s", 0)
art_duration = artifact["end_s"] - artifact["onset_s"]
padding = max(WINDOW_PADDING_S, art_duration * 0.3)
win_start = max(0, artifact["onset_s"] - padding)
win_end = min(duration, artifact["end_s"] + padding)
win_end = min(win_end, win_start + MAX_WINDOW_S)
if win_end - win_start < MIN_WINDOW_S:
win_end = min(win_start + MIN_WINDOW_S, duration)
if not selected_channels:
selected_channels = _get_relevant_channels(artifact_type, artifact.get("channel", ""), all_channels, anns)
channel_indices = [i for i, name in enumerate(all_channels) if name in selected_channels]
if not channel_indices:
channel_indices = list(range(min(8, len(all_channels))))
selected_channels = [all_channels[i] for i in channel_indices]
start_sample = int(win_start * sfreq)
end_sample = int(win_end * sfreq)
try:
signal = read_signal_window(canonical_uri, start_sample, end_sample, channel_indices)
ch_names = [all_channels[i] for i in channel_indices]
except Exception as e:
return (
build_eeg_figure(np.zeros((1, 100)), ["S3 Error"], 256.0, title=str(e)[:80]),
f'{e}
',
gr.CheckboxGroup(choices=all_channels, value=selected_channels),
)
source_uri = info.get("archival_uri", "")
window_anns = get_annotations_in_window(canonical_uri, win_start, win_end, source_uri=source_uri)
fig = build_eeg_figure(
signal, ch_names, sfreq,
start_time_s=win_start,
annotations=window_anns,
title=f"{artifact_type} | {info.get('subject', '?')} | {win_start:.1f}–{win_end:.1f}s",
)
return (
fig,
build_annotation_summary(window_anns),
gr.CheckboxGroup(choices=all_channels, value=selected_channels),
)
def _get_relevant_channels(artifact_type: str, art_channel: str, all_channels: list[str], anns: list[dict]) -> list[str]:
relevant = [a for a in anns if a["label"] == artifact_type]
ann_channels = set(a.get("channel", "") for a in relevant if a.get("channel"))
matched = []
for name in all_channels:
name_clean = name.upper().replace("EEG ", "").replace("-REF", "").replace("-", "").replace(" ", "")
for ann_ch in ann_channels:
parts = ann_ch.upper().replace("-", "")
if parts in name_clean or name_clean in parts:
matched.append(name)
break
if matched:
return list(dict.fromkeys(matched))[:12]
channel_map = {
"eye": ["FP1", "FP2", "F7", "F8", "F3", "F4"],
"muscle": ["T3", "T4", "T5", "T6", "F7", "F8"],
"chew": ["T3", "T4", "T5", "T6", "F7", "F8"],
}
target = next((v for k, v in channel_map.items() if k in artifact_type.lower()),
["FP1", "FP2", "F3", "F4", "C3", "C4", "P3", "P4", "O1", "O2"])
result = [name for name in all_channels if any(t in name.upper() for t in target)]
return result[:12] if result else all_channels[:8]
CSS = """
.gradio-container {max-width: 1600px !important;}
footer {display: none !important;}
"""
with gr.Blocks(title="TUAR EEG Artifact Explorer") as app:
gr.Markdown(
"# TUAR EEG Artifact Explorer\n"
"Browse EEG artifacts by type. Select artifact > recording > instance. "
"Everything streams from S3."
)
with gr.Accordion(
"AWS Credentials" + (" (configured)" if check_aws_credentials() else " (required)"),
open=not check_aws_credentials(),
):
with gr.Row():
aws_key = gr.Textbox(label="Access Key ID", type="password", placeholder="AKIA...", scale=2)
aws_secret = gr.Textbox(label="Secret Access Key", type="password", scale=2)
aws_region = gr.Textbox(label="Region", value="us-east-1", scale=1)
save_btn = gr.Button("Save Credentials", variant="secondary", size="sm")
creds_status = gr.HTML("")
save_btn.click(fn=save_credentials, inputs=[aws_key, aws_secret, aws_region], outputs=[creds_status])
gr.Markdown("---")
status_html = gr.HTML('Loading TUAR recordings…
')
gr.Markdown("### Step 1: Select artifact type")
artifact_type_dropdown = gr.Dropdown(
choices=ALL_ARTIFACT_TYPES, value=None,
label="What artifact are you looking for?", interactive=True,
)
scan_status = gr.HTML("")
gr.Markdown("### Step 2: Select recording")
recording_dropdown = gr.Dropdown(
choices=[], label="Recordings containing this artifact", interactive=True,
)
gr.Markdown("### Step 3: Select specific artifact instance")
instance_dropdown = gr.Dropdown(
choices=[], label="Artifact instances in this recording", interactive=True,
)
gr.Markdown("---")
with gr.Row():
with gr.Column(scale=4):
eeg_plot = gr.Plot(label="EEG Signal")
with gr.Column(scale=1):
gr.Markdown("### Recording Info")
metadata_html = gr.HTML('No recording loaded.
')
gr.Markdown("### Artifacts in View")
annotation_html = gr.HTML('')
gr.Markdown("### Legend")
gr.HTML(build_artifact_legend())
with gr.Accordion("Channel Selection (auto-selected, or pick manually)", open=False):
channel_selector = gr.CheckboxGroup(choices=[], value=[], label="Channels")
# --- Events ---
app.load(fn=init_recordings, inputs=[], outputs=[status_html])
artifact_type_dropdown.change(
fn=on_artifact_type_selected,
inputs=[artifact_type_dropdown],
outputs=[recording_dropdown, scan_status],
)
recording_dropdown.change(
fn=on_recording_selected,
inputs=[artifact_type_dropdown, recording_dropdown],
outputs=[instance_dropdown, metadata_html, eeg_plot, annotation_html, channel_selector],
)
instance_dropdown.change(
fn=on_instance_selected,
inputs=[artifact_type_dropdown, recording_dropdown, instance_dropdown, channel_selector],
outputs=[eeg_plot, annotation_html, channel_selector],
)
channel_selector.change(
fn=on_instance_selected,
inputs=[artifact_type_dropdown, recording_dropdown, instance_dropdown, channel_selector],
outputs=[eeg_plot, annotation_html, channel_selector],
)
if __name__ == "__main__":
port = int(os.environ.get("PORT", 7860))
app.launch(
server_name="0.0.0.0",
server_port=port,
share=False,
theme=gr.themes.Base(primary_hue="blue", secondary_hue="slate", neutral_hue="slate"),
css=CSS,
)