#!/usr/bin/env python3
"""tcp_socket_demo.py

A practical Python 3 TCP socket demo for local network communication.

Features:
- Run as a server or client from one script
- Safe default loopback-only behavior
- Basic validation and clear console output
- Optional custom host, port, and message

Examples:
  # Terminal 1
  python tcp_socket_demo.py server --host 127.0.0.1 --port 65432

  # Terminal 2
  python tcp_socket_demo.py client --host 127.0.0.1 --port 65432 --message "hello from client"

Notes:
- Uses TCP sockets and UTF-8 text encoding.
- Demonstrates a single request/response exchange.
- Intended for local testing, learning, and internal tooling.
"""

from __future__ import annotations

import argparse
import socket
import sys
from dataclasses import dataclass
from typing import Optional

DEFAULT_HOST = "127.0.0.1"
DEFAULT_PORT = 65432
DEFAULT_BUFFER_SIZE = 1024
DEFAULT_MESSAGE = "hello from client"
DEFAULT_TIMEOUT = 5.0


@dataclass(frozen=True)
class Config:
    host: str
    port: int
    buffer_size: int
    message: str
    timeout: float


def parse_args(argv: Optional[list[str]] = None) -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Run a simple TCP server/client demo for local network communication."
    )
    subparsers = parser.add_subparsers(dest="mode", required=True)

    server_parser = subparsers.add_parser("server", help="Start the TCP server")
    add_common_args(server_parser, include_message=False)

    client_parser = subparsers.add_parser("client", help="Start the TCP client")
    add_common_args(client_parser, include_message=True)

    return parser.parse_args(argv)


def add_common_args(parser: argparse.ArgumentParser, include_message: bool) -> None:
    parser.add_argument("--host", default=DEFAULT_HOST, help=f"Bind/connect host (default: {DEFAULT_HOST})")
    parser.add_argument("--port", type=validate_port, default=DEFAULT_PORT, help=f"TCP port (default: {DEFAULT_PORT})")
    parser.add_argument(
        "--buffer-size",
        type=validate_positive_int,
        default=DEFAULT_BUFFER_SIZE,
        help=f"Receive buffer size in bytes (default: {DEFAULT_BUFFER_SIZE})",
    )
    parser.add_argument(
        "--timeout",
        type=validate_non_negative_float,
        default=DEFAULT_TIMEOUT,
        help=f"Socket timeout in seconds (default: {DEFAULT_TIMEOUT})",
    )
    if include_message:
        parser.add_argument(
            "--message",
            default=DEFAULT_MESSAGE,
            help=f"Message to send from the client (default: {DEFAULT_MESSAGE!r})",
        )


def validate_port(value: str | int) -> int:
    try:
        port = int(value)
    except (TypeError, ValueError) as exc:
        raise argparse.ArgumentTypeError("port must be an integer") from exc
    if not (1 <= port <= 65535):
        raise argparse.ArgumentTypeError("port must be between 1 and 65535")
    return port


def validate_positive_int(value: str | int) -> int:
    try:
        number = int(value)
    except (TypeError, ValueError) as exc:
        raise argparse.ArgumentTypeError("value must be an integer") from exc
    if number <= 0:
        raise argparse.ArgumentTypeError("value must be greater than 0")
    return number


def validate_non_negative_float(value: str | float) -> float:
    try:
        number = float(value)
    except (TypeError, ValueError) as exc:
        raise argparse.ArgumentTypeError("timeout must be a number") from exc
    if number < 0:
        raise argparse.ArgumentTypeError("timeout must be 0 or greater")
    return number


def make_config(args: argparse.Namespace) -> Config:
    message = getattr(args, "message", DEFAULT_MESSAGE)
    if not isinstance(message, str) or not message.strip():
        raise ValueError("message must be a non-empty string")
    return Config(
        host=args.host,
        port=args.port,
        buffer_size=args.buffer_size,
        message=message,
        timeout=args.timeout,
    )


def run_server(config: Config) -> int:
    try:
        with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as server_socket:
            server_socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
            server_socket.settimeout(config.timeout)
            server_socket.bind((config.host, config.port))
            server_socket.listen(1)
            print(f"Server listening on {config.host}:{config.port}")

            try:
                conn, addr = server_socket.accept()
            except socket.timeout:
                print(f"No client connected within {config.timeout} seconds")
                return 1

            with conn:
                conn.settimeout(config.timeout)
                print(f"Connected by {addr}")

                data = conn.recv(config.buffer_size)
                if not data:
                    print("No data received")
                    return 1

                try:
                    message = data.decode("utf-8")
                except UnicodeDecodeError:
                    print("Received non-UTF-8 data")
                    return 1

                print(f"Received: {message}")
                response = f"ACK: {message}"
                conn.sendall(response.encode("utf-8"))
                print("Response sent")
                return 0
    except OSError as exc:
        print(f"Server error: {exc}", file=sys.stderr)
        return 1


def run_client(config: Config) -> int:
    try:
        with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as client_socket:
            client_socket.settimeout(config.timeout)
            client_socket.connect((config.host, config.port))
            print(f"Connected to {config.host}:{config.port}")

            client_socket.sendall(config.message.encode("utf-8"))
            print(f"Sent: {config.message}")

            data = client_socket.recv(config.buffer_size)
            if not data:
                print("No response received")
                return 1

            try:
                response = data.decode("utf-8")
            except UnicodeDecodeError:
                print("Received non-UTF-8 response")
                return 1

            print(f"Received: {response}")
            return 0
    except ConnectionRefusedError:
        print("Connection refused: ensure the server is running and reachable", file=sys.stderr)
        return 1
    except socket.timeout:
        print(f"Timed out after {config.timeout} seconds", file=sys.stderr)
        return 1
    except OSError as exc:
        print(f"Client error: {exc}", file=sys.stderr)
        return 1


def main(argv: Optional[list[str]] = None) -> int:
    args = parse_args(argv)
    try:
        config = make_config(args)
    except ValueError as exc:
        print(f"Configuration error: {exc}", file=sys.stderr)
        return 2

    if args.mode == "server":
        return run_server(config)
    if args.mode == "client":
        return run_client(config)

    print("Unknown mode", file=sys.stderr)
    return 2


if __name__ == "__main__":
    raise SystemExit(main())