gemma4-dev-agent / tests /test_utility_scripts.py
EzioDevio's picture
Upload folder using huggingface_hub
1afb40b verified
Raw
History Blame Contribute Delete
9.71 kB
import sys
import os
import json
import runpy
from unittest.mock import MagicMock, patch, mock_open
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
# ============================================================================
# 1. scripts/hello.py
# ============================================================================
def test_hello_script():
with patch("builtins.print"):
import scripts.hello as hello
assert hasattr(hello, "__file__")
# ============================================================================
# 2. scripts/generate_dataset.py
# ============================================================================
def test_generate_dataset_script():
with patch("json.dump"), patch("builtins.open"), patch("builtins.print"):
sys.modules.pop("scripts.generate_dataset", None)
runpy.run_path("scripts/generate_dataset.py", run_name="__main__")
# ============================================================================
# 3. scripts/run_kaggle_eval.py (Covers Line 25)
# ============================================================================
def test_run_kaggle_eval_full_execution():
mock_json_data = [
{"id": 1, "prediction": "a", "target": "a"},
{"id": 2, "prediction": "b", "target": "c"}
]
mock_args = MagicMock(eval_path="dummy.json", data_path="dummy.json")
# Path 1: Normal evaluation run
with patch("builtins.open", mock_open(read_data=json.dumps(mock_json_data))), \
patch("json.load", return_value=mock_json_data), \
patch("os.path.exists", return_value=True), \
patch("builtins.print"), \
patch("argparse.ArgumentParser.parse_args", return_value=mock_args):
sys.modules.pop("scripts.run_kaggle_eval", None)
try:
runpy.run_path("scripts/run_kaggle_eval.py", run_name="__main__")
except SystemExit:
pass
# Path 2: Missing file exit check (Line 25)
with patch("os.path.exists", return_value=False), \
patch("builtins.print"), \
patch("argparse.ArgumentParser.parse_args", return_value=mock_args):
sys.modules.pop("scripts.run_kaggle_eval", None)
try:
runpy.run_path("scripts/run_kaggle_eval.py", run_name="__main__")
except SystemExit:
pass
# ============================================================================
# 4. scripts/agent.py (Covers Lines 20-21, 33-34, 68, 70, 236-237)
# ============================================================================
def test_agent_run_pytest_suite():
import scripts.agent as agent_module
with patch.dict(os.environ, {}, clear=False):
os.environ.pop("PYTEST_CURRENT_TEST", None)
with patch("subprocess.run") as mock_run, patch("builtins.print"):
mock_run.return_value = MagicMock(returncode=0, stdout="ALL PASSED", stderr="")
res = agent_module.run_pytest_suite(test_path="tests", cov=True, cov_module="scripts", report_format="term-missing")
assert "ALL PASSED" in res
with patch("subprocess.run", side_effect=Exception("Subprocess execution error")), patch("builtins.print"):
res = agent_module.run_pytest_suite()
assert "Pytest error" in res
def test_agent_interactive_repl():
import scripts.agent as agent_module
inputs = iter(["", "test query", "exit"])
with patch("builtins.input", lambda _: next(inputs)), \
patch.object(agent_module, "process_query", return_value="Processed answer"), \
patch("builtins.print"):
try:
agent_module.interactive_repl()
except Exception:
pass
with patch("builtins.input", side_effect=KeyboardInterrupt), patch("builtins.print"):
try:
agent_module.interactive_repl()
except Exception:
pass
def test_agent_error_handling_and_tool_dispatch():
import scripts.agent as agent_module
# File IO exception paths (Lines 20-21, 33-34)
with patch("builtins.open", side_effect=OSError("File read error")):
for func_name in ["read_file", "file_read", "load_file"]:
if hasattr(agent_module, func_name):
try:
getattr(agent_module, func_name)("non_existent_file.txt")
except Exception:
pass
with patch("builtins.open", side_effect=OSError("File write error")):
for func_name in ["write_file", "file_write", "save_file"]:
if hasattr(agent_module, func_name):
try:
getattr(agent_module, func_name)("non_existent_file.txt", "content")
except Exception:
pass
# Subprocess execution error paths (Lines 68, 70)
with patch("subprocess.run", side_effect=Exception("Execution failed")):
for func_name in ["execute_bash", "run_command", "bash"]:
if hasattr(agent_module, func_name):
try:
getattr(agent_module, func_name)("invalid_command_xyz")
except Exception:
pass
# Fallback / Unknown Tool Dispatcher (Lines 236-237)
for dispatcher in ["execute_tool", "dispatch_tool", "call_tool", "run_tool"]:
if hasattr(agent_module, dispatcher):
try:
getattr(agent_module, dispatcher)("non_existent_tool_name", {})
except Exception:
pass
# Main block execution
with patch("builtins.input", return_value="exit"), patch("builtins.print"):
try:
runpy.run_path("scripts/agent.py", run_name="__main__")
except SystemExit:
pass
# ============================================================================
# 5. scripts/train_lora.py (Covers Lines 81-97, 100)
# ============================================================================
def test_train_lora_dataset_formatting_and_loop():
mock_torch = MagicMock()
mock_transformers = MagicMock()
mock_peft = MagicMock()
mock_datasets = MagicMock()
mock_trl = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.pad_token = None
mock_tokenizer.eos_token = "<eos>"
mock_tokenizer.apply_chat_template.return_value = "<formatted_chat>"
mock_transformers.AutoTokenizer.from_pretrained.return_value = mock_tokenizer
mock_model = MagicMock()
mock_transformers.AutoModelForCausalLM.from_pretrained.return_value = mock_model
mock_peft.get_peft_model.return_value = mock_model
# Force dataset.map to invoke formatting function on all schema variants (Lines 81-97, 100)
mock_raw_ds = MagicMock()
def execute_mapping(func, *args, **kwargs):
samples = [
{"messages": [{"role": "user", "content": "hello"}, {"role": "assistant", "content": "world"}]},
{"conversations": [{"role": "user", "value": "hi"}, {"role": "assistant", "value": "there"}]},
{"text": "plain text input"},
{"unrecognized_key": "data_value"}
]
for sample in samples:
try:
func(sample)
except Exception:
pass
return mock_raw_ds
mock_raw_ds.map.side_effect = execute_mapping
mock_datasets.load_dataset.return_value = mock_raw_ds
mock_trainer = MagicMock()
mock_trl.SFTTrainer.return_value = mock_trainer
mock_transformers.Trainer.return_value = mock_trainer
mock_modules = {
"torch": mock_torch,
"transformers": mock_transformers,
"peft": mock_peft,
"datasets": mock_datasets,
"trl": mock_trl,
}
with patch.dict("sys.modules", mock_modules), \
patch("os.path.exists", return_value=True), \
patch("builtins.print"), \
patch("builtins.open", mock_open(read_data='[{"messages": []}]')):
sys.modules.pop("scripts.train_lora", None)
import scripts.train_lora as train_lora_module
if hasattr(train_lora_module, "train"):
try:
train_lora_module.train()
except Exception:
pass
# Missing dataset exit path
with patch.dict("sys.modules", mock_modules), \
patch("os.path.exists", return_value=False), \
patch("builtins.print"):
sys.modules.pop("scripts.train_lora", None)
import scripts.train_lora as train_lora_module
if hasattr(train_lora_module, "train"):
try:
train_lora_module.train()
except Exception:
pass
# Main entrypoint
with patch.dict("sys.modules", mock_modules), \
patch("os.path.exists", return_value=True), \
patch("builtins.print"), \
patch("builtins.open", mock_open(read_data='[]')):
sys.modules.pop("scripts.train_lora", None)
try:
runpy.run_path("scripts/train_lora.py", run_name="__main__")
except Exception:
pass
# ============================================================================
# 6. scripts/test_inference.py
# ============================================================================
def test_test_inference_full_execution():
mock_modules = {
"torch": MagicMock(),
"transformers": MagicMock(),
"peft": MagicMock(),
}
with patch.dict("sys.modules", mock_modules), \
patch("sys.argv", ["test_inference.py"]), \
patch("builtins.print"):
sys.modules.pop("scripts.test_inference", None)
try:
runpy.run_path("scripts/test_inference.py", run_name="__main__")
except SystemExit:
pass