"""
JX-SMS Python SDK v1.0

Send SMS messages via the JX-SMS API from any Python application.
Requires: requests (pip install requests)

Usage:
    from jxsms import JxSms

    sms = JxSms('YOUR_API_KEY', 'YOUR_API_SECRET')
    result = sms.send('256770123456', 'Hello from JX-SMS!')
    print(result)

Docs: https://jxsms.com/api-docs
"""

import re
import requests


class JxSmsError(Exception):
    """Raised when the JX-SMS API returns an error."""

    def __init__(self, message, status_code=None, response=None):
        super().__init__(message)
        self.status_code = status_code
        self.response = response


class JxSms:
    """JX-SMS API client."""

    def __init__(self, api_key, api_secret, base_url='https://jxsms.com/api/v1', timeout=30):
        """
        Args:
            api_key:    Your 32-character API key
            api_secret: Your 48-character API secret
            base_url:   API base URL (override for local dev)
            timeout:    Request timeout in seconds
        """
        self.api_key = api_key
        self.api_secret = api_secret
        self.base_url = base_url.rstrip('/')
        self.timeout = timeout
        self._session = requests.Session()
        self._session.headers.update({
            'X-API-Key': self.api_key,
            'X-API-Secret': self.api_secret,
            'Accept': 'application/json',
        })

    def send(self, phone_number, message, sender_id='JX-SMS'):
        """
        Send a single SMS.

        Args:
            phone_number: Recipient (e.g. '256770123456' or '0770123456')
            message:      SMS text (max 320 chars, >160 = 2 SMS)
            sender_id:    Sender name (max 11 chars)

        Returns:
            dict: {success, message_id, phone, telecom, status, sms_count, cost, balance}
        """
        return self._post('/send', {
            'phone_number': self.normalize_phone(phone_number),
            'message': message,
            'sender_id': sender_id,
        })

    def send_bulk(self, phone_numbers, message, sender_id='JX-SMS'):
        """
        Send bulk SMS to multiple recipients.

        Args:
            phone_numbers: List of numbers or comma-separated string
            message:       SMS text
            sender_id:     Sender name

        Returns:
            dict: {success, campaign_id, total, sent, queued, cost, balance}
        """
        if isinstance(phone_numbers, (list, tuple)):
            phone_numbers = ','.join(self.normalize_phone(n) for n in phone_numbers)

        return self._post('/send-bulk', {
            'phone_numbers': phone_numbers,
            'message': message,
            'sender_id': sender_id,
        })

    def balance(self):
        """
        Check account balance.

        Returns:
            dict: {success, balance, currency}
        """
        return self._get('/balance')

    def status(self, message_id):
        """
        Check delivery status of a sent message.

        Args:
            message_id: The message_id from the send response

        Returns:
            dict: {success, phone, status, telecom, cost, sent_at, delivered_at}
        """
        return self._get('/status', {'message_id': message_id})

    def contacts(self, group_id=None):
        """
        List contact groups, or contacts within a group.

        Args:
            group_id: If None, returns groups. If set, returns contacts in that group.

        Returns:
            dict: {success, groups: [...]} or {success, contacts: [...]}
        """
        params = {}
        if group_id is not None:
            params['group_id'] = group_id
        return self._get('/contacts', params)

    def pricing(self):
        """
        Get your SMS pricing per network.

        Returns:
            dict: {success, pricing: [{prefix, telecom, price_per_sms, currency}]}
        """
        return self._get('/pricing')

    # ─── Internal ────────────────────────────────────────────

    @staticmethod
    def normalize_phone(number):
        """Normalize Ugandan phone number to 256XXXXXXXXX format."""
        number = re.sub(r'[^0-9]', '', str(number))
        if len(number) == 10 and number[0] == '0':
            number = '256' + number[1:]
        elif len(number) == 9:
            number = '256' + number
        return number

    def _post(self, endpoint, data):
        return self._request('POST', endpoint, data=data)

    def _get(self, endpoint, params=None):
        return self._request('GET', endpoint, params=params)

    def _request(self, method, endpoint, data=None, params=None):
        url = self.base_url + endpoint

        try:
            resp = self._session.request(
                method, url,
                data=data,
                params=params,
                timeout=self.timeout,
            )
        except requests.ConnectionError as e:
            raise JxSmsError(f'Connection error: {e}')
        except requests.Timeout:
            raise JxSmsError('Request timed out')

        try:
            result = resp.json()
        except ValueError:
            raise JxSmsError(f'Invalid response from API (HTTP {resp.status_code})')

        if result.get('error'):
            raise JxSmsError(
                result.get('message', 'Unknown API error'),
                status_code=resp.status_code,
                response=result,
            )

        return result

    def __repr__(self):
        masked = self.api_key[:8] + '...' if self.api_key else 'None'
        return f'JxSms(api_key={masked}, base_url={self.base_url})'
