import sys
import numpy as np

def calculate_average_distance(file_path):
    with open(file_path, 'r') as f:
        lines = f.readlines()

    distances = []
    for line in lines:
        # Skip lines that are not part of the matrix (e.g., headers)
        if line.strip() == "" or line.strip().startswith("Distance") or line.strip().startswith("Using"):
            continue
        
        # Split each line to extract distance values
        values = line.strip().split()
        
        # Skip the first value which is likely the row/column label
        distances.extend([float(x) for x in values[:-1] if x.replace('.', '', 1).replace('-', '', 1).isdigit()])

    # Calculate the average distance, excluding zeroes (distances of a sequence to itself)
    distances = [dist for dist in distances if dist > 0]
    
    if len(distances) == 0:
        return None
    
    average_distance = np.mean(distances)
    return average_distance

if __name__ == "__main__":
    if len(sys.argv) != 2:
        print("Usage: python calculate_intracluster_distance.py <distance_matrix_file>")
        sys.exit(1)
    
    distance_matrix_file = sys.argv[1]
    average_distance = calculate_average_distance(distance_matrix_file)
    
    if average_distance is not None:
        # Print the file name and average distance in a tab-separated format
        print(f"{distance_matrix_file}\t{average_distance:.4f}")
    else:
        print(f"{distance_matrix_file}\tFailed to calculate the average pairwise distance.")

