<%
    import pandas as pd
    def plot_embdedding(data, plotter_list):
        sns = plotter_list["sns"]
        ax = sns.scatterplot(data=data[data['group_seed'] == 'other'], x='coord1', y='coord2', 
            color='lightgrey', legend=False, zorder=1, s=5, marker='+', alpha=0.4, linewidth=0.6)

        if 'group' in data.columns:
            ax = sns.scatterplot(data=data[data['group'] != 'other'], x='coord1', y='coord2', hue='group', palette='Set2', 
                color='lightgrey', legend=True, zorder=2, s=10, edgecolor="black", linewidth=0.6)

        ax = sns.scatterplot(data=df[df['group_seed'] != 'other'], x='coord1', y='coord2',
         hue='group_seed', palette='pastel', color='white', legend=False, zorder=3, s=30, marker='D', edgecolor="black", linewidth=0.6)
    
        handles, labels = ax.get_legend_handles_labels()
        sorted_labels = sorted(labels)
        new_order = [ labels.index(lab) for lab in sorted_labels]
        ax.legend([handles[idx] for idx in new_order], sorted_labels)
        #ax.get_legend().remove() # Commented for single plot
        #ax.legend(handles, labels)
        sns.move_legend(ax, "lower center", bbox_to_anchor=(0.5, 1), title=None, ncols=8, columnspacing=0.1, handletextpad=0.03, frameon=False)
        #ax.legend(handles, labels, ncols=8, columnspacing=0.1, handletextpad=0.03, loc='lower center', bbox_to_anchor=(.5, 1), frameon=False)
        #plotter_list["plt"].tight_layout() # Commented for single plot
        #plotter_list["plt"].subplots_adjust(right=0.7)  # Make space on the right # Commented for single plot

    
    node2seed = {}
    for seed, nodes in plotter.hash_vars["seeds2explore"].items():
        for node in nodes:
            if not node2seed.get(node): node2seed[node] = seed

    node2group = {}
    if plotter.hash_vars["groups2explore"]:
        for group, nodes in plotter.hash_vars["groups2explore"].items():
            for node in nodes:
                if not node2group.get(node): node2group[node] = group
%>

## Seeds LCC.
% if plotter.hash_vars.get("seeds2lcc"):
${plotter.create_title("LCC by seed", id='net_study', hlevel=1, indexable=True, clickable=False)}
<div style="overflow: hidden; display: flex; flex-direction: col; justify-content: center;">
<%
table_lcc = [["seed_name","net","lcc"]]
for seed, net_lcc in plotter.hash_vars["seeds2lcc"].items():
    for net, lcc in net_lcc.items():
        table_lcc.append([seed, net, lcc])
plotter.hash_vars["table_lcc"] = table_lcc
%>
${plotter.barplot(id='table_lcc', responsive= False, header=True,
                         fields = [0,2],
                         x_label = 'LCC size',
                         height = '400px', width= '400px',
                         smp_attr = [1], title="",
                         config = {
                                'showLegend' : True,
                                'colorBy' : 'seed_name',
                                'setMinX': 0,
                                "titleFontStyle": "italic",
                                "titleScaleFontFactor": 0.7,
                                "smpLabelScaleFontFactor": 1,
                                "axisTickScaleFontFactor": 0.7,
                                "segregateSamplesBy": "net"
                                })}
</div>
%endif
## Seeds subgraphs.
% if plotter.hash_vars.get("seeds2subgraph"):
    <% plotter.set_header() %>
    ${plotter.create_title("Network study on every group of genes", id='net_study', hlevel=1, indexable=True, clickable=False)}
    % for seed, net2subgraph in plotter.hash_vars["seeds2subgraph"].items():
        ${plotter.create_title(f"Group of genes: {seed}", id='net_study', hlevel=2, indexable=True, clickable=False)}
        % for net, subgraph in net2subgraph.items():
            ${plotter.create_title(f"<b>{net}</b>", id='net_study', hlevel=3, indexable=True, clickable=False)}
            <% 
            key = net+"_"+seed
            plotter.hash_vars[key] = {"graph": subgraph, 
                "layers": ["layer", "layer"], 
                "reference_nodes": plotter.hash_vars["seeds2explore"][seed], 
                "group_nodes": plotter.hash_vars["groups2explore"] if plotter.hash_vars.get("groups2explore") else {}}  
            %>
            ${plotter.network(id = key, **plotter.hash_vars["plot_options"])} 
        %endfor
    %endfor
%endif

% if plotter.hash_vars["net2embedding_proj"]:
    ${plotter.create_title("Network embedding and projection", id='net_emb', hlevel=1, indexable=True, clickable=False)}
    % for graph, embedding in plotter.hash_vars["net2embedding_proj"].items():
        ${plotter.create_title(f"embedding in {graph}", id=f"net_emb_{graph}", hlevel=2, indexable=True, clickable=False)}
        <% 
        coords, nodes = embedding 
        df = pd.DataFrame(coords, columns= ["coord1","coord2"])
        df = df.assign(group_seed=[node2seed[node] if node2seed.get(node) else "other" for node in nodes]) 

        if plotter.hash_vars["groups2explore"]:
            df = df.assign(group=[node2group[node] if node2group.get(node) else "other" for node in nodes])
            df.sort_values(by='group', inplace=True) # Sort by group to assign always the same color to the same group 

        plotter.hash_vars["net_proj"] = df
        %>
        ${plotter.static_plot_main(id="net_proj", raw= True, plotting_function= lambda data, plotter_list: plot_embdedding(data, plotter_list), theme='fast', width=800, height=800, dpi=150)} 
    % endfor
% endif

## Net sims
% if plotter.hash_vars.get("net_sims"):
${plotter.create_title("Network similarity edge", id='net_edge', hlevel=1, indexable=True, clickable=False)}
<%
sim_nets, network_ids = plotter.hash_vars["net_sims"]
table_sims = [["",*network_ids]]
for i, net_id in enumerate(network_ids):
    table_sims.append([net_id,*sim_nets[i,:].tolist()])
plotter.hash_vars["table_sims"] = table_sims
%>
${plotter.corplot(id = 'table_sims', header = True, row_names = True, correlationAxis = 'variables', config= {"correlationType":"circle"})}
%endif
