import sqlite3
import os
import hashlib
import datetime
import json

class DatabaseManager:
    def __init__(self, db_path):
        self.db_path = db_path
        self.conn = None
        self.ensure_db_exists()
        
    def ensure_db_exists(self):
        """Create the database and schema if it doesn't exist"""
        db_dir = os.path.dirname(self.db_path)
        if not os.path.exists(db_dir):
            os.makedirs(db_dir)
            
        if not os.path.exists(self.db_path):
            self.connect()
            with open('database/schema.sql', 'r') as f:
                self.conn.executescript(f.read())
            self.conn.commit()
            self.create_admin_user('admin', 'password')  # Default admin user
    
    def connect(self):
        """Connect to the SQLite database"""
        if self.conn is None:
            self.conn = sqlite3.connect(self.db_path)
            self.conn.row_factory = sqlite3.Row
        return self.conn
    
    def close(self):
        """Close the database connection"""
        if self.conn:
            self.conn.close()
            self.conn = None
    
    def create_admin_user(self, username, password):
        """Create an admin user with the given credentials"""
        password_hash = hashlib.sha256(password.encode()).hexdigest()
        cursor = self.conn.cursor()
        cursor.execute(
            "INSERT INTO users (username, password_hash) VALUES (?, ?)",
            (username, password_hash)
        )
        self.conn.commit()
    
    def authenticate_user(self, username, password):
        """Authenticate a user and return user data if successful"""
        password_hash = hashlib.sha256(password.encode()).hexdigest()
        cursor = self.conn.cursor()
        cursor.execute(
            "SELECT id, username FROM users WHERE username = ? AND password_hash = ?",
            (username, password_hash)
        )
        user = cursor.fetchone()
        
        if user:
            # Update last login time
            cursor.execute(
                "UPDATE users SET last_login = ? WHERE id = ?",
                (datetime.datetime.now().isoformat(), user['id'])
            )
            self.conn.commit()
            return dict(user)
        return None
    
    def add_artist(self, name):
        """Add an artist to the database or get existing one"""
        cursor = self.conn.cursor()
        cursor.execute("SELECT id FROM artists WHERE name = ?", (name,))
        artist = cursor.fetchone()
        
        if artist:
            return artist['id']
        
        cursor.execute("INSERT INTO artists (name) VALUES (?)", (name,))
        self.conn.commit()
        return cursor.lastrowid
    
    def add_album(self, title, artist_id, release_year=None, cover_path=None):
        """Add an album to the database or get existing one"""
        cursor = self.conn.cursor()
        cursor.execute(
            "SELECT id FROM albums WHERE title = ? AND artist_id = ?", 
            (title, artist_id)
        )
        album = cursor.fetchone()
        
        if album:
            return album['id']
        
        cursor.execute(
            "INSERT INTO albums (title, artist_id, release_year, cover_path) VALUES (?, ?, ?, ?)",
            (title, artist_id, release_year, cover_path)
        )
        self.conn.commit()
        return cursor.lastrowid
    
    def add_track(self, track_data):
        """Add a track to the database"""
        cursor = self.conn.cursor()
        
        # Check if track already exists
        cursor.execute("SELECT id FROM tracks WHERE file_path = ?", (track_data['file_path'],))
        track = cursor.fetchone()
        
        if track:
            return track['id']
        
        # Add the track
        cursor.execute("""
            INSERT INTO tracks (
                title, artist_id, album_id, genre, duration, 
                track_number, file_path, file_format, bitrate
            ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
        """, (
            track_data['title'],
            track_data['artist_id'],
            track_data['album_id'],
            track_data.get('genre'),
            track_data.get('duration'),
            track_data.get('track_number'),
            track_data['file_path'],
            track_data.get('file_format'),
            track_data.get('bitrate')
        ))
        self.conn.commit()
        return cursor.lastrowid
    
    def search_tracks(self, query, limit=50):
        """Search for tracks by title, artist, or album"""
        search_term = f"%{query}%"
        cursor = self.conn.cursor()
        cursor.execute("""
            SELECT t.id, t.title, t.file_path, t.duration, t.track_number, 
                   a.name as artist_name, al.title as album_title, al.cover_path
            FROM tracks t
            LEFT JOIN artists a ON t.artist_id = a.id
            LEFT JOIN albums al ON t.album_id = al.id
            WHERE t.title LIKE ? OR a.name LIKE ? OR al.title LIKE ?
            LIMIT ?
        """, (search_term, search_term, search_term, limit))
        
        tracks = [dict(row) for row in cursor.fetchall()]
        return tracks
    
    def get_albums(self, limit=50, offset=0):
        """Get albums with their cover art"""
        cursor = self.conn.cursor()
        cursor.execute("""
            SELECT al.id, al.title, al.release_year, al.cover_path, a.name as artist_name,
                   COUNT(t.id) as track_count
            FROM albums al
            LEFT JOIN artists a ON al.artist_id = a.id
            LEFT JOIN tracks t ON t.album_id = al.id
            GROUP BY al.id
            ORDER BY al.title
            LIMIT ? OFFSET ?
        """, (limit, offset))
        
        albums = [dict(row) for row in cursor.fetchall()]
        return albums
    
    def get_album_tracks(self, album_id):
        """Get all tracks for a specific album"""
        cursor = self.conn.cursor()
        cursor.execute("""
            SELECT t.id, t.title, t.file_path, t.duration, t.track_number, 
                   a.name as artist_name
            FROM tracks t
            LEFT JOIN artists a ON t.artist_id = a.id
            WHERE t.album_id = ?
            ORDER BY t.track_number, t.title
        """, (album_id,))
        
        tracks = [dict(row) for row in cursor.fetchall()]
        return tracks
    
    def get_playlists(self, user_id):
        """Get all playlists for a user"""
        cursor = self.conn.cursor()
        cursor.execute("""
            SELECT p.id, p.name, p.created_at, COUNT(pt.track_id) as track_count
            FROM playlists p
            LEFT JOIN playlist_tracks pt ON p.id = pt.playlist_id
            WHERE p.user_id = ?
            GROUP BY p.id
            ORDER BY p.updated_at DESC
        """, (user_id,))
        
        playlists = [dict(row) for row in cursor.fetchall()]
        return playlists
    
    def create_playlist(self, user_id, name):
        """Create a new playlist"""
        cursor = self.conn.cursor()
        now = datetime.datetime.now().isoformat()
        cursor.execute(
            "INSERT INTO playlists (name, user_id, created_at, updated_at) VALUES (?, ?, ?, ?)",
            (name, user_id, now, now)
        )
        self.conn.commit()
        return cursor.lastrowid
    
    def add_track_to_playlist(self, playlist_id, track_id, position=None):
        """Add a track to a playlist"""
        cursor = self.conn.cursor()
        
        # Get the next position if not specified
        if position is None:
            cursor.execute(
                "SELECT COALESCE(MAX(position), 0) + 1 FROM playlist_tracks WHERE playlist_id = ?",
                (playlist_id,)
            )
            position = cursor.fetchone()[0]
        
        cursor.execute(
            "INSERT OR REPLACE INTO playlist_tracks (playlist_id, track_id, position) VALUES (?, ?, ?)",
            (playlist_id, track_id, position)
        )
        
        # Update the playlist's updated_at timestamp
        cursor.execute(
            "UPDATE playlists SET updated_at = ? WHERE id = ?",
            (datetime.datetime.now().isoformat(), playlist_id)
        )
        
        self.conn.commit()
    
    def get_playlist_tracks(self, playlist_id):
        """Get all tracks in a playlist"""
        cursor = self.conn.cursor()
        cursor.execute("""
            SELECT t.id, t.title, t.file_path, t.duration, 
                   a.name as artist_name, al.title as album_title, al.cover_path,
                   pt.position
            FROM playlist_tracks pt
            JOIN tracks t ON pt.track_id = t.id
            LEFT JOIN artists a ON t.artist_id = a.id
            LEFT JOIN albums al ON t.album_id = al.id
            WHERE pt.playlist_id = ?
            ORDER BY pt.position
        """, (playlist_id,))
        
        tracks = [dict(row) for row in cursor.fetchall()]
        return tracks
    
    def add_to_favorites(self, user_id, track_id):
        """Add a track to user's favorites"""
        cursor = self.conn.cursor()
        cursor.execute(
            "INSERT OR IGNORE INTO favorites (user_id, track_id) VALUES (?, ?)",
            (user_id, track_id)
        )
        self.conn.commit()
    
    def remove_from_favorites(self, user_id, track_id):
        """Remove a track from user's favorites"""
        cursor = self.conn.cursor()
        cursor.execute(
            "DELETE FROM favorites WHERE user_id = ? AND track_id = ?",
            (user_id, track_id)
        )
        self.conn.commit()
    
    def get_favorites(self, user_id):
        """Get all favorite tracks for a user"""
        cursor = self.conn.cursor()
        cursor.execute("""
            SELECT t.id, t.title, t.file_path, t.duration, 
                   a.name as artist_name, al.title as album_title, al.cover_path
            FROM favorites f
            JOIN tracks t ON f.track_id = t.id
            LEFT JOIN artists a ON t.artist_id = a.id
            LEFT JOIN albums al ON t.album_id = al.id
            WHERE f.user_id = ?
            ORDER BY f.added_at DESC
        """, (user_id,))
        
        tracks = [dict(row) for row in cursor.fetchall()]
        return tracks
    
    def update_play_count(self, track_id):
        """Increment the play count for a track"""
        cursor = self.conn.cursor()
        cursor.execute(
            "UPDATE tracks SET play_count = play_count + 1 WHERE id = ?",
            (track_id,)
        )
        self.conn.commit()
    
    def get_stats(self):
        """Get database statistics"""
        cursor = self.conn.cursor()
        cursor.execute("SELECT COUNT(*) FROM tracks")
        track_count = cursor.fetchone()[0]
        
        cursor.execute("SELECT COUNT(*) FROM albums")
        album_count = cursor.fetchone()[0]
        
        cursor.execute("SELECT COUNT(*) FROM artists")
        artist_count = cursor.fetchone()[0]
        
        cursor.execute("SELECT COUNT(*) FROM playlists")
        playlist_count = cursor.fetchone()[0]
        
        return {
            'tracks': track_count,
            'albums': album_count,
            'artists': artist_count,
            'playlists': playlist_count
        }