import os
import re
from collections import defaultdict
from docx import Document
import spacy
import networkx as nx
import plotly.graph_objects as go

# Load spaCy's English NLP model
try:
    nlp = spacy.load('en_core_web_sm')
    print("Model loaded successfully!")
except Exception as e:
    print("An error occurred while loading the model:", e)

# Specify the directory where your files are stored
directory_path = r'C:\Users\eek205\OneDrive - University of Exeter\Desktop\Hashem'

# List all Word documents in the directory
doc_files = [file for file in os.listdir(directory_path) if file.endswith('.docx')]

# Function to extract year from the file title
def extract_year_from_title(filename):
    match = re.search(r'\d{4}', filename)
    if match:
        return int(match.group(0))
    return None

# Extract years from filenames and group files by decade
decade_groups = defaultdict(list)
for file in doc_files:
    year = extract_year_from_title(file)
    if year:
        decade = (year // 10) * 10
        decade_groups[decade].append(file)

# Read and preprocess texts
def read_docx(file_path):
    doc = Document(file_path)
    return " ".join([para.text for para in doc.paragraphs])

def preprocess(text):
    doc = nlp(text)
    return [token.lemma_.lower() for token in doc if token.is_alpha and not token.is_stop]

# Process texts by decade
texts_by_decade = defaultdict(list)
for decade, files in decade_groups.items():
    for file in files:
        file_path = os.path.join(directory_path, file)
        text = read_docx(file_path)
        processed_text = preprocess(text)
        texts_by_decade[decade].extend(processed_text)

# Build co-occurrence networks for each decade
def build_cooccurrence_network(tokens, window_size=4):
    G = nx.Graph()
    for i in range(len(tokens) - window_size + 1):
        window = tokens[i:i + window_size]
        for j in range(len(window)):
            for k in range(j + 1, len(window)):
                if window[j] != window[k]:
                    if G.has_edge(window[j], window[k]):
                        G[window[j]][window[k]]['weight'] += 1
                    else:
                        G.add_edge(window[j], window[k], weight=1)
    return G

networks_by_decade = {decade: build_cooccurrence_network(tokens) for decade, tokens in texts_by_decade.items()}

# Function to visualize the network using Plotly
def plotly_network(G, title="Network Graph"):
    pos = nx.spring_layout(G)  # Calculate layout positions
    centrality = nx.degree_centrality(G)
    sorted_nodes = sorted(G.nodes(), key=lambda node: centrality[node], reverse=True)
    pivotal_nodes = set(sorted_nodes[:int(0.1 * len(sorted_nodes))])  # Top 10%

    edge_x = []
    edge_y = []
    node_x = []
    node_y = []
    node_text = []
    node_color = []

    for edge in G.edges():
        x0, y0 = pos[edge[0]]
        x1, y1 = pos[edge[1]]
        edge_x.extend([x0, x1, None])
        edge_y.extend([y0, y1, None])

    for node in G.nodes():
        x, y = pos[node]
        node_x.append(x)
        node_y.append(y)
        if node in pivotal_nodes:
            node_text.append(f"<b>{node}</b>")
            node_color.append('red')
        else:
            node_text.append("")
            node_color.append('blue')

    node_trace = go.Scatter(
        x=node_x, y=node_y,
        mode='markers+text',
        marker=dict(size=[centrality[n] * 100 for n in G.nodes()], color=node_color),
        text=node_text,
        textposition="top center",
        textfont=dict(color='red', size=12)
    )

    edge_trace = go.Scatter(
        x=edge_x, y=edge_y,
        line=dict(width=0.5, color='grey'),
        hoverinfo='none',
        mode='lines'
    )

    fig = go.Figure(data=[edge_trace, node_trace], layout=go.Layout(
        showlegend=False,
        hovermode='closest',
        title=title,
        title_x=0.5,
        margin=dict(b=0, l=0, r=0, t=40),
        xaxis=dict(showgrid=False, zeroline=False, showticklabels=False),
        yaxis=dict(showgrid=False, zeroline=False, showticklabels=False)
        )
    )
    fig.show()

# Visualize networks for each decade
for decade, network in networks_by_decade.items():
    plotly_network(network, title=f"Interactive Network Graph for the {decade}s")

def save_networks_to_graphml(networks, directory):
    os.makedirs(directory, exist_ok=True)  # Ensure the directory exists
    for decade, network in networks.items():
        # Filename format: Network_<decade>.graphml
        path = os.path.join(directory, f'Network_{decade}.graphml')
        nx.write_graphml(network, path)
        print(f"Saved network for the {decade}s to {path}")

# Specify the directory where you want to save the GraphML files
output_directory = r'C:\Users\eek205\OneDrive - University of Exeter\Desktop\GraphML'

# Save all decade networks to GraphML files
save_networks_to_graphml(networks_by_decade, output_directory)

