#!/usr/bin/env zsh

# Function to install packages required by a single Python file
function install_requirements_for_file() {
    local python_file=$1
    shift # Remove the first argument

    # Use Python to extract required packages and install them
    python3 -c "
import ast, sys

def is_standard_lib(module_name):
    return module_name in getattr(sys, 'stdlib_module_names', [])

def get_top_level_module(name):
    return name.split('.')[0]

with open('$python_file', 'r') as file:
    root = ast.parse(file.read(), filename='$python_file')

    imports = {
        get_top_level_module(alias.name)
        for node in ast.walk(root) if isinstance(node, ast.Import)
        for alias in node.names if not is_standard_lib(get_top_level_module(alias.name))
    }

    from_imports = {
        get_top_level_module(node.module)
        for node in ast.walk(root) if isinstance(node, ast.ImportFrom)
        if not is_standard_lib(get_top_level_module(node.module))
    }

for package in imports.union(from_imports):
    print(package)
" | xargs pip install "$@"
}

# Main script execution
install_requirements_for_file "$@"
