from RFEM.initModel import Model, Calculate_all
from RFEM.baseSettings import BaseSettings
from RFEM.BasicObjects.node import Node
from RFEM.BasicObjects.line import Line
from RFEM.BasicObjects.thickness import Thickness
from RFEM.BasicObjects.material import Material
from RFEM.BasicObjects.surface import Surface
from RFEM.TypesForLines.lineSupport import LineSupport
from RFEM.enums import LineSupportType, NodalLoadDirection, GlobalAxesOrientationType
from RFEM.Loads.nodalLoad import NodalLoad
from RFEM.LoadCasesAndCombinations.loadCase import LoadCase
from RFEM.Calculate.meshSettings import MeshSettings
from RFEM.Results.meshTables import MeshTables
import pandas as pd
import plotly.graph_objs as go

Model(True, "steelBeamSurfaces.rf6")
Model.clientModel.service.delete_all()
BaseSettings(global_axes_orientation=GlobalAxesOrientationType.E_GLOBAL_AXES_ORIENTATION_ZUP)
Node(1, 0,0,0)
Node(2, 0,0,0.35)
Node(3, 0,-0.1,0)
Node(4, 0,0.1,0)
Node(5, 0,-0.1,0.35)
Node(6, 0,0.1,0.35)
Node(7, 2,0,0)
Node(8, 2,0,0.15)
Node(9, 2.000,	-0.100,	0.000)
Node(10, 2.000,	0.100,	0.000)
Node(11, 2.000,	-0.100,	0.150)
Node(12, 2.000,	0.100,	0.150)

Line(1, "1 2")
Line(2, "3 1")
Line(3, "1 4")
Line(4, "5 2")
Line(5, "2 6")
Line(6, "7 8")
Line(7, "9 7")
Line(8, "7 10")
Line(9, "11 8")
Line(10, "8 12")
Line(11, "1 7")
Line(12, "2 8")
Line(13, "3 9")
Line(14, "4 10")
Line(15, "5 11")
Line(16, "6 12")

Material(1, "S235")
Thickness(1, "web", 1, 0.01)
Thickness(2, "flange", 1, 0.016)

Surface(1, "11 1 12 6", 1)
Surface(2, "3 11 8 14", 2)
Surface(3, "2 13 7 11", 2)
Surface(4, "5 16 10 12", 2)
Surface(5, "4 12 9 15", 2)

LineSupport(1, '1', LineSupportType.FIXED)

LoadCase(1)

NodalLoad(1, 1, '7', NodalLoadDirection.LOAD_DIRECTION_GLOBAL_Z_OR_USER_DEFINED_W, -10000000)
NodalLoad(2, 1, '7', NodalLoadDirection.LOAD_DIRECTION_GLOBAL_Y_OR_USER_DEFINED_V, 20000)

meshLength = 0.5

MeshSettings(commonConfig={'general_target_length_of_fe': meshLength})

Calculate_all()

allSurfaces = MeshTables.GetAllFE2DElements()
surfaces = pd.DataFrame(allSurfaces)
surfaces = surfaces[['surface_no','FE_node1_no', 'FE_node2_no', 'FE_node3_no', 'FE_node4_no']].astype(int)

allFENodes = MeshTables.GetAllFENodes()

meshNodes = pd.DataFrame(allFENodes)

meshNodes = meshNodes[['x', 'y', 'z']]

allFeNodesDeformed = MeshTables.GetAllFENodesDeformed()

deformedMeshNodes = pd.DataFrame(allFeNodesDeformed)

deformedMeshNodes = deformedMeshNodes[['x', 'y', 'z']]

def plot(orig= None):
    vertices = {
        'x': orig['x'],  
        'y': orig['y'],  
        'z': orig['z']
    }

    faces=[]
    for i in range(len(surfaces)):
        faces.append((surfaces['FE_node1_no'][i]-1, surfaces['FE_node2_no'][i]-1, surfaces['FE_node3_no'][i]-1, surfaces['FE_node4_no'][i]-1))

    triangles = []
    for quad in faces:
        triangles.extend([quad[0], quad[1], quad[2]])
        triangles.extend([quad[2], quad[3], quad[0]])

    mesh = go.Mesh3d(
        x=vertices['x'],
        y=vertices['y'],
        z=vertices['z'],
        i=[triangles[i] for i in range(0, len(triangles), 3)],
        j=[triangles[i+1] for i in range(0, len(triangles), 3)],
        k=[triangles[i+2] for i in range(0, len(triangles), 3)],
        opacity=0.5,
        color='cyan'
    )

    edge_traces = []

    def create_edge_trace(v1, v2):
        return go.Scatter3d(
            x=[vertices['x'][v1], vertices['x'][v2], None], 
            y=[vertices['y'][v1], vertices['y'][v2], None], 
            z=[vertices['z'][v1], vertices['z'][v2], None], 
            mode='lines',
            line=dict(color='white', width=2),
            hoverinfo='none'
        )

    for quad in faces:
        edge_traces.append(create_edge_trace(quad[0], quad[1]))
        edge_traces.append(create_edge_trace(quad[1], quad[2]))
        edge_traces.append(create_edge_trace(quad[2], quad[3]))
        edge_traces.append(create_edge_trace(quad[3], quad[0]))

    layout = go.Layout(
        plot_bgcolor='black',
        paper_bgcolor='black',
        showlegend=False,
        scene=dict(
            xaxis=dict(showbackground=False, visible=False, range= [-2,2]),
            yaxis=dict(showbackground=False, visible=False, range= [-2,2]),
            zaxis=dict(showbackground=False, visible=False, range= [-2,2]),
            aspectratio=dict(x=2, y=2, z=2)
        ),
        margin=dict(l=0, r=0, b=0, t=0),
        
    )

    fig = go.Figure(data=[mesh]+edge_traces, layout=layout)
    camera = dict(
    up=dict(x=0, y=0, z=1),
        center=dict(x=0, y=0, z=0),
        eye=dict(x=0.5, y=-2, z=0.5)
    )

    fig.update_layout(scene_camera= camera)

    return fig

print(meshNodes)
print(deformedMeshNodes)



fig = plot(meshNodes)
deformedFig = plot(deformedMeshNodes)
fig.show()
deformedFig.show()