import sshtunnel
from pprint import pprint
from db import DB
from sqlalchemy.sql import text
from typing import TypedDict
from typing import Generator
from pathlib import Path
import csv
import sys

from alumnos import do_action, READER_OPTS

type StudentId = dict[str, int]

DBURL = 'mysql+pymysql://moodleuser:moodleMysqlV1wild@localhost:6606/moodle'
COURSE = 'PROGI25'


def get_course_id(course_shortname: str, conn: DB) -> int:
    sql = text(
        """
        SELECT id FROM mdl_course
        WHERE shortname = :shortname
        """
    )
    result = conn.execute_dict(sql, shortname=course_shortname)
    row = result.first()
    if row is None:
        raise ValueError(f"Course with shortname '{course_shortname}' not found.")
    return row['id']

def get_students_ids(courseid: int, conn: DB) -> StudentId:
    sql = text(
        """
SELECT user.id as id, user.email AS email, course.fullname AS Course
FROM mdl_course AS course
JOIN mdl_enrol AS en ON en.courseid = course.id
JOIN mdl_user_enrolments AS ue ON ue.enrolid = en.id
JOIN mdl_user AS user ON ue.userid = user.id
WHERE course.id=:courseid
        """
    )
    result = conn.execute_dict(sql, courseid=courseid)
    return {row['email']: row['id'] for row in result}



def save_students_ids(students: StudentId,
                      path_ori: Path) -> None:

    path_dest = path_ori.with_name(path_ori.stem + '_moodle_ids.csv')
    with open(path_dest, 'w') as f, \
         open(path_ori, 'r') as f_ori:
        reader = csv.DictReader(f_ori, **READER_OPTS)
        fieldnames = reader.fieldnames + ['MOODLE_ID']
        writer = csv.DictWriter(f, fieldnames=fieldnames, **READER_OPTS)
        writer.writeheader()
        for st in reader:
            email = st['CORREO']
            st['MOODLE_ID'] = students.get(email, '')
            writer.writerow(st)

def print_students(course: str) -> None:
    with sshtunnel.open_tunnel(
            'wild.ucm.es',
            ssh_username='root',
            ssh_pkey='~/.ssh/id_rsa',
            remote_bind_address=('localhost', 3306),
            local_bind_address=('localhost', 6606)) as tunnel, \
        DB(DBURL) as conn:

        courseid = get_course_id(course, conn)
        print(f"Course ID for '{course}': {courseid}")
        for email, student in get_students_ids(courseid, conn).items():
            print(email, student)

def main():
    if len(sys.argv) < 2:
        print(f'Usage: {sys.argv[0]} <filename>')
        sys.exit(1)

    filename = Path(sys.argv[1])
    with sshtunnel.open_tunnel(
            'wild.ucm.es',
            ssh_username='root',
            ssh_pkey='~/.ssh/id_rsa',
            remote_bind_address=('localhost', 3306),
            local_bind_address=('localhost', 6606)) as tunnel, \
        DB(DBURL) as conn:

        courseid = get_course_id(COURSE, conn)
        print(f"Course ID for '{COURSE}': {courseid}")
        students = get_students_ids(courseid, conn)
        save_students_ids(students, filename)

if __name__ == "__main__":
    main()
