/*
 * Simplified EPICS Channel Access client.
 *
 * This is an independent implementation of the CA wire protocol. It supports
 * IPv4 discovery and scalar DBR_STRING, SHORT, FLOAT, ENUM, CHAR, LONG, and
 * DOUBLE transfers. It deliberately omits arrays, subscriptions, preemptive
 * callbacks, the CA repeater, and automatic reconnection.
 *
 * Generated with OpenAI Codex 5.6 Sol-high, 29.08.2026, S. Ritt
 */

#include "epics_ca.h"

#include <arpa/inet.h>
#include <errno.h>
#include <fcntl.h>
#include <ifaddrs.h>
#include <net/if.h>
#include <netdb.h>
#include <netinet/in.h>
#include <netinet/tcp.h>
#include <poll.h>
#include <pwd.h>
#include <sys/socket.h>
#include <sys/types.h>
#include <unistd.h>

#include <algorithm>
#include <chrono>
#include <cstdint>
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <limits>
#include <memory>
#include <mutex>
#include <new>
#include <string>
#include <unordered_map>
#include <vector>

struct ca_channel {
   uint32_t magic;
};

namespace {

using Clock = std::chrono::steady_clock;

const uint32_t kChannelMagic = 0x43414844; // "CAHD"
const uint16_t kDefaultServerPort = 5064;
const uint16_t kProtocolRevision = 13;
const size_t kHeaderSize = 16;
const size_t kMaxStringSize = 40;
const size_t kMaxPayloadSize = 16 * 1024 * 1024;

enum Command : uint16_t {
   kVersion = 0,
   kWrite = 4,
   kSearch = 6,
   kError = 11,
   kReadNotify = 15,
   kCreateChannel = 18,
   kClientName = 20,
   kHostName = 21,
   kAccessRights = 22,
   kEcho = 23,
   kCreateChannelFail = 26,
   kServerDisconnect = 27
};

enum class ChannelState {
   Searching,
   Creating,
   Connected,
   Failed,
   Disconnected
};

enum class CircuitState {
   Connecting,
   Connected,
   Dead
};

struct Circuit;

struct Channel {
   ca_channel handle = {kChannelMagic};
   std::string name;
   uint32_t cid = 0;
   uint32_t sid = 0;
   uint16_t native_type = 0;
   uint32_t native_count = 0;
   uint32_t access_rights = 0;
   bool access_known = false;
   bool pend_connect = true;
   ChannelState state = ChannelState::Searching;
   Circuit *circuit = nullptr;
   Clock::time_point next_search = Clock::now();
   std::chrono::milliseconds search_delay{30};
};

struct Circuit {
   sockaddr_in address{};
   int fd = -1;
   CircuitState state = CircuitState::Connecting;
   std::vector<uint8_t> input;
   std::vector<uint8_t> output;
   size_t output_offset = 0;
   std::vector<Channel *> channels;
   Clock::time_point last_activity = Clock::now();
};

struct PendingGet {
   Channel *channel = nullptr;
   chtype type = DBR_STRING;
   void *destination = nullptr;
};

struct Header {
   uint16_t command = 0;
   uint32_t payload_size = 0;
   uint16_t data_type = 0;
   uint32_t data_count = 0;
   uint32_t parameter1 = 0;
   uint32_t parameter2 = 0;
   size_t header_size = kHeaderSize;
};

uint16_t read_u16(const uint8_t *data)
{
   uint16_t value;
   std::memcpy(&value, data, sizeof(value));
   return ntohs(value);
}

uint32_t read_u32(const uint8_t *data)
{
   uint32_t value;
   std::memcpy(&value, data, sizeof(value));
   return ntohl(value);
}

uint64_t read_u64(const uint8_t *data)
{
   return (static_cast<uint64_t>(read_u32(data)) << 32) | read_u32(data + 4);
}

void append_u16(std::vector<uint8_t> &buffer, uint16_t value)
{
   value = htons(value);
   const uint8_t *bytes = reinterpret_cast<const uint8_t *>(&value);
   buffer.insert(buffer.end(), bytes, bytes + sizeof(value));
}

void append_u32(std::vector<uint8_t> &buffer, uint32_t value)
{
   value = htonl(value);
   const uint8_t *bytes = reinterpret_cast<const uint8_t *>(&value);
   buffer.insert(buffer.end(), bytes, bytes + sizeof(value));
}

void append_u64(std::vector<uint8_t> &buffer, uint64_t value)
{
   append_u32(buffer, static_cast<uint32_t>(value >> 32));
   append_u32(buffer, static_cast<uint32_t>(value));
}

size_t aligned_size(size_t size)
{
   return (size + 7) & ~static_cast<size_t>(7);
}

void append_message(std::vector<uint8_t> &buffer, uint16_t command,
                    uint16_t data_type, uint32_t data_count,
                    uint32_t parameter1, uint32_t parameter2,
                    const uint8_t *payload = nullptr, size_t payload_size = 0)
{
   const size_t padded_size = aligned_size(payload_size);
   if (padded_size > std::numeric_limits<uint16_t>::max())
      return;

   append_u16(buffer, command);
   append_u16(buffer, static_cast<uint16_t>(padded_size));
   append_u16(buffer, data_type);
   append_u16(buffer, static_cast<uint16_t>(data_count));
   append_u32(buffer, parameter1);
   append_u32(buffer, parameter2);
   if (payload_size)
      buffer.insert(buffer.end(), payload, payload + payload_size);
   buffer.insert(buffer.end(), padded_size - payload_size, 0);
}

bool decode_header(const uint8_t *data, size_t size, Header &header)
{
   if (size < kHeaderSize)
      return false;

   header.command = read_u16(data);
   header.payload_size = read_u16(data + 2);
   header.data_type = read_u16(data + 4);
   header.data_count = read_u16(data + 6);
   header.parameter1 = read_u32(data + 8);
   header.parameter2 = read_u32(data + 12);
   header.header_size = kHeaderSize;

   if (header.payload_size == 0xffff && header.data_count == 0) {
      if (size < 24)
         return false;
      header.payload_size = read_u32(data + 16);
      header.data_count = read_u32(data + 20);
      header.header_size = 24;
   }

   return true;
}

bool set_nonblocking(int fd)
{
   int flags = fcntl(fd, F_GETFL, 0);
   return flags >= 0 && fcntl(fd, F_SETFL, flags | O_NONBLOCK) == 0;
}

bool same_address(const sockaddr_in &left, const sockaddr_in &right)
{
   return left.sin_addr.s_addr == right.sin_addr.s_addr &&
          left.sin_port == right.sin_port;
}

class ClientContext {
public:
   int initialize()
   {
      if (initialized_)
         return ECA_NORMAL;

      server_port_ = parse_port(std::getenv("EPICS_CA_SERVER_PORT"),
                                kDefaultServerPort);
      configure_destinations();

      udp_fd_ = socket(AF_INET, SOCK_DGRAM, 0);
      if (udp_fd_ < 0)
         return ECA_SOCK;

      int enabled = 1;
      setsockopt(udp_fd_, SOL_SOCKET, SO_BROADCAST, &enabled, sizeof(enabled));

      sockaddr_in local{};
      local.sin_family = AF_INET;
      local.sin_addr.s_addr = htonl(INADDR_ANY);
      local.sin_port = 0;
      if (bind(udp_fd_, reinterpret_cast<sockaddr *>(&local), sizeof(local)) < 0 ||
          !set_nonblocking(udp_fd_)) {
         close(udp_fd_);
         udp_fd_ = -1;
         return ECA_SOCK;
      }

      initialized_ = true;
      return ECA_NORMAL;
   }

   void reset()
   {
      for (auto &circuit : circuits_)
         close_circuit(*circuit);
      if (udp_fd_ >= 0)
         close(udp_fd_);

      udp_fd_ = -1;
      initialized_ = false;
      destinations_.clear();
      pending_gets_.clear();
      channel_index_.clear();
      channels_.clear();
      circuits_.clear();
      next_id_ = 1;
      next_sequence_ = 1;
   }

   int create_channel(const char *name, caCh *connection_handler,
                      capri priority, chid *result)
   {
      if (!result)
         return ECA_BADCHID;
      *result = nullptr;
      if (!name || !name[0] || std::strlen(name) > 500)
         return ECA_BADSTR;
      if (connection_handler)
         return ECA_NOSUPPORT;
      if (priority > 99)
         return ECA_BADPRIORITY;
      if (destinations_.empty())
         return ECA_NOSEARCHADDR;

      try {
         std::unique_ptr<Channel> channel(new Channel);
         channel->name = name;
         channel->cid = allocate_id();
         channel->next_search = Clock::now();
         Channel *raw = channel.get();
         chid handle = &raw->handle;
         channels_.push_back(std::move(channel));
         channel_index_[handle] = raw;
         *result = handle;
      } catch (const std::bad_alloc &) {
         return ECA_ALLOCMEM;
      }
      return ECA_NORMAL;
   }

   int get(chtype type, chid handle, void *destination)
   {
      Channel *channel = find_channel(handle);
      if (!channel)
         return ECA_BADCHID;
      if (!valid_type(type))
         return ECA_BADTYPE;
      if (!destination)
         return ECA_GETFAIL;
      if (channel->state != ChannelState::Connected || !channel->circuit ||
          channel->circuit->state != CircuitState::Connected)
         return ECA_DISCONN;
      if (channel->access_known && !(channel->access_rights & 1))
         return ECA_NORDACCESS;

      uint32_t ioid = allocate_id();
      try {
         PendingGet pending;
         pending.channel = channel;
         pending.type = type;
         pending.destination = destination;
         pending_gets_[ioid] = pending;
         append_message(channel->circuit->output, kReadNotify,
                        static_cast<uint16_t>(type), 1, channel->sid, ioid);
      } catch (const std::bad_alloc &) {
         pending_gets_.erase(ioid);
         return ECA_ALLOCMEM;
      }

      if (!flush_circuit(*channel->circuit)) {
         pending_gets_.erase(ioid);
         return ECA_DISCONN;
      }
      return ECA_NORMAL;
   }

   int put(chtype type, chid handle, const void *value)
   {
      Channel *channel = find_channel(handle);
      if (!channel)
         return ECA_BADCHID;
      if (!valid_type(type))
         return ECA_BADTYPE;
      if (!value)
         return ECA_PUTFAIL;
      if (channel->state != ChannelState::Connected || !channel->circuit ||
          channel->circuit->state != CircuitState::Connected)
         return ECA_DISCONN;
      if (channel->access_known && !(channel->access_rights & 2))
         return ECA_NOWTACCESS;

      std::vector<uint8_t> payload;
      int status = encode_value(type, value, payload);
      if (status != ECA_NORMAL)
         return status;

      try {
         append_message(channel->circuit->output, kWrite,
                        static_cast<uint16_t>(type), 1, channel->sid,
                        channel->cid, payload.data(), payload.size());
      } catch (const std::bad_alloc &) {
         return ECA_ALLOCMEM;
      }

      return flush_circuit(*channel->circuit) ? ECA_NORMAL : ECA_DISCONN;
   }

   int pend_io(ca_real timeout)
   {
      if (timeout < 0)
         return ECA_TIMEOUT;

      for (auto &circuit : circuits_)
         flush_circuit(*circuit);

      if (!has_outstanding())
         return ECA_NORMAL;

      const bool infinite = timeout == 0;
      const Clock::time_point deadline = infinite
         ? Clock::time_point::max()
         : Clock::now() + std::chrono::duration_cast<Clock::duration>(
                             std::chrono::duration<double>(timeout));

      while (has_outstanding()) {
         Clock::time_point now = Clock::now();
         if (!infinite && now >= deadline) {
            expire_pending();
            return ECA_TIMEOUT;
         }

         send_due_searches(now);
         std::vector<pollfd> poll_fds;
         std::vector<Circuit *> poll_circuits;
         build_poll_list(poll_fds, poll_circuits);

         int wait_ms = calculate_wait(now, deadline, infinite);
         int status = poll(poll_fds.data(), poll_fds.size(), wait_ms);
         if (status < 0) {
            if (errno == EINTR)
               continue;
            expire_pending();
            return ECA_INTERNAL;
         }

         size_t index = 0;
         if (udp_fd_ >= 0) {
            if (poll_fds[index].revents & POLLIN)
               receive_udp();
            ++index;
         }

         for (Circuit *circuit : poll_circuits) {
            short events = poll_fds[index++].revents;
            if (!events)
               continue;
            if (circuit->state == CircuitState::Connecting &&
                (events & (POLLOUT | POLLERR | POLLHUP))) {
               finish_connect(*circuit);
            } else if (circuit->state == CircuitState::Connected) {
               if (events & POLLOUT)
                  flush_circuit(*circuit);
               if (events & POLLIN)
                  receive_tcp(*circuit);
               if (circuit->state == CircuitState::Connected &&
                   (events & (POLLERR | POLLHUP | POLLNVAL)))
                  disconnect_circuit(*circuit);
            }
         }
      }

      for (auto &circuit : circuits_)
         flush_circuit(*circuit);
      return ECA_NORMAL;
   }

private:
   static uint16_t parse_port(const char *text, uint16_t fallback)
   {
      if (!text || !text[0])
         return fallback;
      char *end = nullptr;
      long value = std::strtol(text, &end, 10);
      if (!end || *end || value < 1 || value > 65535)
         return fallback;
      return static_cast<uint16_t>(value);
   }

   static bool valid_type(chtype type)
   {
      return type >= DBR_STRING && type <= DBR_DOUBLE;
   }

   uint32_t allocate_id()
   {
      uint32_t result = next_id_++;
      if (next_id_ == 0)
         next_id_ = 1;
      return result;
   }

   Channel *find_channel(chid handle)
   {
      auto found = channel_index_.find(handle);
      if (found == channel_index_.end() || !handle ||
          handle->magic != kChannelMagic)
         return nullptr;
      return found->second;
   }

   void add_destination(const sockaddr_in &address)
   {
      auto duplicate = std::find_if(destinations_.begin(), destinations_.end(),
         [&](const sockaddr_in &item) { return same_address(item, address); });
      if (duplicate == destinations_.end())
         destinations_.push_back(address);
   }

   void add_host_destination(const std::string &token)
   {
      std::string host = token;
      uint16_t port = server_port_;
      size_t colon = token.rfind(':');
      if (colon != std::string::npos && token.find(':') == colon) {
         uint16_t parsed = parse_port(token.c_str() + colon + 1, 0);
         if (parsed) {
            host = token.substr(0, colon);
            port = parsed;
         }
      }
      if (host.empty())
         return;

      addrinfo hints{};
      hints.ai_family = AF_INET;
      hints.ai_socktype = SOCK_DGRAM;
      addrinfo *addresses = nullptr;
      if (getaddrinfo(host.c_str(), nullptr, &hints, &addresses) != 0)
         return;

      for (addrinfo *item = addresses; item; item = item->ai_next) {
         if (item->ai_addrlen < sizeof(sockaddr_in))
            continue;
         sockaddr_in address = *reinterpret_cast<sockaddr_in *>(item->ai_addr);
         address.sin_port = htons(port);
         add_destination(address);
      }
      freeaddrinfo(addresses);
   }

   void add_auto_destinations()
   {
      ifaddrs *interfaces = nullptr;
      if (getifaddrs(&interfaces) != 0)
         return;

      for (ifaddrs *item = interfaces; item; item = item->ifa_next) {
         if (!item->ifa_addr || item->ifa_addr->sa_family != AF_INET ||
             !(item->ifa_flags & IFF_UP) || !(item->ifa_flags & IFF_BROADCAST) ||
             !item->ifa_broadaddr)
            continue;
         sockaddr_in address =
            *reinterpret_cast<sockaddr_in *>(item->ifa_broadaddr);
         address.sin_port = htons(server_port_);
         add_destination(address);
      }
      freeifaddrs(interfaces);

      if (destinations_.empty()) {
         sockaddr_in loopback{};
         loopback.sin_family = AF_INET;
         loopback.sin_addr.s_addr = htonl(INADDR_LOOPBACK);
         loopback.sin_port = htons(server_port_);
         add_destination(loopback);
      }
   }

   void configure_destinations()
   {
      const char *automatic = std::getenv("EPICS_CA_AUTO_ADDR_LIST");
      bool use_automatic = true;
      if (automatic) {
         std::string value(automatic);
         use_automatic = value.find("NO") == std::string::npos &&
                         value.find("no") == std::string::npos;
      }
      if (use_automatic)
         add_auto_destinations();

      const char *list = std::getenv("EPICS_CA_ADDR_LIST");
      if (!list)
         return;
      const char *begin = list;
      while (*begin) {
         while (*begin == ' ' || *begin == '\t' || *begin == '\r' ||
                *begin == '\n')
            ++begin;
         const char *end = begin;
         while (*end && *end != ' ' && *end != '\t' && *end != '\r' &&
                *end != '\n')
            ++end;
         if (end != begin)
            add_host_destination(std::string(begin, end));
         begin = end;
      }
   }

   int encode_value(chtype type, const void *value,
                    std::vector<uint8_t> &payload)
   {
      try {
         switch (type) {
         case DBR_STRING: {
            const char *text = static_cast<const char *>(value);
            size_t length = strnlen(text, kMaxStringSize);
            if (length == kMaxStringSize)
               return ECA_BADCOUNT;
            payload.insert(payload.end(), text, text + length + 1);
            break;
         }
         case DBR_SHORT:
         case DBR_ENUM: {
            uint16_t integer;
            std::memcpy(&integer, value, sizeof(integer));
            append_u16(payload, integer);
            break;
         }
         case DBR_FLOAT: {
            uint32_t bits;
            std::memcpy(&bits, value, sizeof(bits));
            append_u32(payload, bits);
            break;
         }
         case DBR_CHAR:
            payload.push_back(*static_cast<const uint8_t *>(value));
            break;
         case DBR_LONG: {
            uint32_t integer;
            std::memcpy(&integer, value, sizeof(integer));
            append_u32(payload, integer);
            break;
         }
         case DBR_DOUBLE: {
            uint64_t bits;
            std::memcpy(&bits, value, sizeof(bits));
            append_u64(payload, bits);
            break;
         }
         default:
            return ECA_BADTYPE;
         }
      } catch (const std::bad_alloc &) {
         return ECA_ALLOCMEM;
      }
      return ECA_NORMAL;
   }

   bool decode_value(chtype type, const uint8_t *payload, size_t payload_size,
                     void *destination)
   {
      switch (type) {
      case DBR_STRING: {
         std::memset(destination, 0, kMaxStringSize);
         size_t length = std::min(payload_size, kMaxStringSize - 1);
         std::memcpy(destination, payload, length);
         return true;
      }
      case DBR_SHORT:
      case DBR_ENUM: {
         if (payload_size < 2)
            return false;
         uint16_t value = read_u16(payload);
         std::memcpy(destination, &value, sizeof(value));
         return true;
      }
      case DBR_FLOAT: {
         if (payload_size < 4)
            return false;
         uint32_t bits = read_u32(payload);
         std::memcpy(destination, &bits, sizeof(bits));
         return true;
      }
      case DBR_CHAR:
         if (!payload_size)
            return false;
         *static_cast<uint8_t *>(destination) = payload[0];
         return true;
      case DBR_LONG: {
         if (payload_size < 4)
            return false;
         uint32_t value = read_u32(payload);
         std::memcpy(destination, &value, sizeof(value));
         return true;
      }
      case DBR_DOUBLE: {
         if (payload_size < 8)
            return false;
         uint64_t bits = read_u64(payload);
         std::memcpy(destination, &bits, sizeof(bits));
         return true;
      }
      default:
         return false;
      }
   }

   bool has_outstanding() const
   {
      if (!pending_gets_.empty())
         return true;
      for (const auto &channel : channels_) {
         if (channel->pend_connect && channel->state != ChannelState::Connected)
            return true;
      }
      return false;
   }

   void expire_pending()
   {
      pending_gets_.clear();
      for (auto &channel : channels_)
         channel->pend_connect = false;
   }

   void send_due_searches(Clock::time_point now)
   {
      for (auto &channel_ptr : channels_) {
         Channel &channel = *channel_ptr;
         if (channel.state != ChannelState::Searching || now < channel.next_search)
            continue;

         std::vector<uint8_t> datagram;
         append_message(datagram, kVersion, 1, kProtocolRevision,
                        next_sequence_++, 0);
         append_message(datagram, kSearch, 5, kProtocolRevision,
                        channel.cid, channel.cid,
                        reinterpret_cast<const uint8_t *>(channel.name.c_str()),
                        channel.name.size() + 1);
         for (const sockaddr_in &destination : destinations_) {
            sendto(udp_fd_, datagram.data(), datagram.size(), 0,
                   reinterpret_cast<const sockaddr *>(&destination),
                   sizeof(destination));
         }

         channel.next_search = now + channel.search_delay;
         channel.search_delay = std::min(channel.search_delay * 2,
                                         std::chrono::milliseconds(1000));
      }
   }

   int calculate_wait(Clock::time_point now, Clock::time_point deadline,
                      bool infinite) const
   {
      Clock::time_point wake = deadline;
      for (const auto &channel : channels_) {
         if (channel->state == ChannelState::Searching)
            wake = std::min(wake, channel->next_search);
      }
      if (infinite && wake == Clock::time_point::max())
         return -1;
      if (wake <= now)
         return 0;
      auto milliseconds =
         std::chrono::duration_cast<std::chrono::milliseconds>(wake - now).count();
      if (milliseconds > std::numeric_limits<int>::max())
         return std::numeric_limits<int>::max();
      return static_cast<int>(std::max<int64_t>(1, milliseconds));
   }

   void build_poll_list(std::vector<pollfd> &poll_fds,
                        std::vector<Circuit *> &poll_circuits)
   {
      if (udp_fd_ >= 0)
         poll_fds.push_back({udp_fd_, POLLIN, 0});
      for (auto &circuit_ptr : circuits_) {
         Circuit &circuit = *circuit_ptr;
         if (circuit.state == CircuitState::Dead || circuit.fd < 0)
            continue;
         short events = POLLIN;
         if (circuit.state == CircuitState::Connecting ||
             circuit.output_offset < circuit.output.size())
            events |= POLLOUT;
         poll_fds.push_back({circuit.fd, events, 0});
         poll_circuits.push_back(&circuit);
      }
   }

   Circuit *find_circuit(const sockaddr_in &address)
   {
      for (auto &circuit : circuits_) {
         if (circuit->state != CircuitState::Dead &&
             same_address(circuit->address, address))
            return circuit.get();
      }
      return nullptr;
   }

   void receive_udp()
   {
      while (true) {
         uint8_t data[65536];
         sockaddr_in source{};
         socklen_t source_size = sizeof(source);
         ssize_t received = recvfrom(udp_fd_, data, sizeof(data), 0,
            reinterpret_cast<sockaddr *>(&source), &source_size);
         if (received < 0) {
            if (errno == EAGAIN || errno == EWOULDBLOCK)
               return;
            return;
         }

         size_t offset = 0;
         while (offset + kHeaderSize <= static_cast<size_t>(received)) {
            Header header;
            if (!decode_header(data + offset, received - offset, header))
               break;
            size_t total = header.header_size + header.payload_size;
            if (total > static_cast<size_t>(received) - offset)
               break;
            if (header.command == kSearch)
               handle_search_reply(header, data + offset + header.header_size,
                                   source);
            offset += total;
         }
      }
   }

   void handle_search_reply(const Header &header, const uint8_t *payload,
                            const sockaddr_in &source)
   {
      auto channel_it = std::find_if(channels_.begin(), channels_.end(),
         [&](const std::unique_ptr<Channel> &item) {
            return item->cid == header.parameter2;
         });
      if (channel_it == channels_.end())
         return;
      Channel &channel = **channel_it;
      if (channel.state != ChannelState::Searching)
         return;

      uint16_t minor_version = 0;
      if (header.payload_size >= 2)
         minor_version = read_u16(payload);

      sockaddr_in server = source;
      if (minor_version >= 8 && header.parameter1 != INADDR_BROADCAST)
         server.sin_addr.s_addr = htonl(header.parameter1);
      if (minor_version >= 5)
         server.sin_port = htons(header.data_type);
      else
         server.sin_port = htons(server_port_);

      Circuit *circuit = find_circuit(server);
      if (!circuit)
         circuit = start_circuit(server);
      if (!circuit)
         return;

      channel.circuit = circuit;
      channel.state = ChannelState::Creating;
      if (std::find(circuit->channels.begin(), circuit->channels.end(), &channel) ==
          circuit->channels.end())
         circuit->channels.push_back(&channel);
      if (circuit->state == CircuitState::Connected) {
         queue_create(channel);
         flush_circuit(*circuit);
      }
   }

   Circuit *start_circuit(const sockaddr_in &address)
   {
      std::unique_ptr<Circuit> circuit(new (std::nothrow) Circuit);
      if (!circuit)
         return nullptr;
      circuit->address = address;
      circuit->fd = socket(AF_INET, SOCK_STREAM, 0);
      if (circuit->fd < 0 || !set_nonblocking(circuit->fd)) {
         if (circuit->fd >= 0)
            close(circuit->fd);
         return nullptr;
      }

      int enabled = 1;
      setsockopt(circuit->fd, IPPROTO_TCP, TCP_NODELAY, &enabled, sizeof(enabled));
      setsockopt(circuit->fd, SOL_SOCKET, SO_KEEPALIVE, &enabled, sizeof(enabled));
#ifdef SO_NOSIGPIPE
      setsockopt(circuit->fd, SOL_SOCKET, SO_NOSIGPIPE, &enabled, sizeof(enabled));
#endif

      Circuit *raw = circuit.get();
      circuits_.push_back(std::move(circuit));
      int status = connect(raw->fd, reinterpret_cast<const sockaddr *>(&address),
                           sizeof(address));
      if (status == 0) {
         activate_circuit(*raw);
      } else if (errno != EINPROGRESS) {
         close_circuit(*raw);
         return nullptr;
      }
      return raw;
   }

   void finish_connect(Circuit &circuit)
   {
      int error = 0;
      socklen_t size = sizeof(error);
      if (getsockopt(circuit.fd, SOL_SOCKET, SO_ERROR, &error, &size) < 0 || error) {
         disconnect_circuit(circuit);
         return;
      }
      activate_circuit(circuit);
   }

   void activate_circuit(Circuit &circuit)
   {
      circuit.state = CircuitState::Connected;
      circuit.last_activity = Clock::now();
      append_message(circuit.output, kVersion, 0, kProtocolRevision, 0, 0);

      char user[128] = "anonymous";
      passwd *account = getpwuid(geteuid());
      if (account && account->pw_name)
         std::snprintf(user, sizeof(user), "%s", account->pw_name);
      append_message(circuit.output, kClientName, 0, 0, 0, 0,
                     reinterpret_cast<const uint8_t *>(user),
                     std::strlen(user) + 1);

      char host[256] = "unknown";
      if (gethostname(host, sizeof(host)) != 0)
         std::snprintf(host, sizeof(host), "unknown");
      host[sizeof(host) - 1] = 0;
      append_message(circuit.output, kHostName, 0, 0, 0, 0,
                     reinterpret_cast<const uint8_t *>(host),
                     std::strlen(host) + 1);

      for (Channel *channel : circuit.channels)
         queue_create(*channel);
      flush_circuit(circuit);
   }

   void queue_create(Channel &channel)
   {
      if (!channel.circuit || channel.circuit->state != CircuitState::Connected)
         return;
      channel.state = ChannelState::Creating;
      append_message(channel.circuit->output, kCreateChannel, 0, 0,
                     channel.cid, kProtocolRevision,
                     reinterpret_cast<const uint8_t *>(channel.name.c_str()),
                     channel.name.size() + 1);
   }

   bool flush_circuit(Circuit &circuit)
   {
      if (circuit.state != CircuitState::Connected)
         return circuit.state == CircuitState::Connecting;
      while (circuit.output_offset < circuit.output.size()) {
#ifdef MSG_NOSIGNAL
         const int send_flags = MSG_NOSIGNAL;
#else
         const int send_flags = 0;
#endif
         ssize_t sent = send(circuit.fd,
            circuit.output.data() + circuit.output_offset,
            circuit.output.size() - circuit.output_offset, send_flags);
         if (sent > 0) {
            circuit.output_offset += static_cast<size_t>(sent);
            circuit.last_activity = Clock::now();
            continue;
         }
         if (sent < 0 && errno == EINTR)
            continue;
         if (sent < 0 && (errno == EAGAIN || errno == EWOULDBLOCK))
            return true;
         disconnect_circuit(circuit);
         return false;
      }
      circuit.output.clear();
      circuit.output_offset = 0;
      return true;
   }

   void receive_tcp(Circuit &circuit)
   {
      uint8_t data[8192];
      bool peer_closed = false;
      while (circuit.state == CircuitState::Connected) {
         ssize_t received = recv(circuit.fd, data, sizeof(data), 0);
         if (received > 0) {
            circuit.input.insert(circuit.input.end(), data, data + received);
            circuit.last_activity = Clock::now();
            continue;
         }
         if (received == 0) {
            peer_closed = true;
            break;
         }
         if (errno == EINTR)
            continue;
         if (errno == EAGAIN || errno == EWOULDBLOCK)
            break;
         disconnect_circuit(circuit);
         return;
      }

      size_t offset = 0;
      while (circuit.state == CircuitState::Connected &&
             circuit.input.size() - offset >= kHeaderSize) {
         Header header;
         if (!decode_header(circuit.input.data() + offset,
                            circuit.input.size() - offset, header))
            break;
         if (header.payload_size > kMaxPayloadSize) {
            disconnect_circuit(circuit);
            return;
         }
         size_t total = header.header_size + header.payload_size;
         if (total > circuit.input.size() - offset)
            break;
         process_tcp_message(circuit, header,
            circuit.input.data() + offset + header.header_size);
         offset += total;
      }
      if (offset)
         circuit.input.erase(circuit.input.begin(), circuit.input.begin() + offset);
      if (peer_closed)
         disconnect_circuit(circuit);
   }

   Channel *find_channel_by_cid(uint32_t cid)
   {
      for (auto &channel : channels_) {
         if (channel->cid == cid)
            return channel.get();
      }
      return nullptr;
   }

   void process_tcp_message(Circuit &circuit, const Header &header,
                            const uint8_t *payload)
   {
      switch (header.command) {
      case kVersion:
         break;
      case kAccessRights: {
         Channel *channel = find_channel_by_cid(header.parameter1);
         if (channel) {
            channel->access_rights = header.parameter2;
            channel->access_known = true;
         }
         break;
      }
      case kCreateChannel: {
         Channel *channel = find_channel_by_cid(header.parameter1);
         if (channel && channel->circuit == &circuit) {
            channel->sid = header.parameter2;
            channel->native_type = header.data_type;
            channel->native_count = header.data_count;
            channel->state = ChannelState::Connected;
            channel->pend_connect = false;
         }
         break;
      }
      case kCreateChannelFail: {
         Channel *channel = find_channel_by_cid(header.parameter1);
         if (channel)
            channel->state = ChannelState::Failed;
         break;
      }
      case kReadNotify: {
         auto found = pending_gets_.find(header.parameter2);
         if (found != pending_gets_.end()) {
            if (header.parameter1 == ECA_NORMAL)
               decode_value(found->second.type, payload, header.payload_size,
                            found->second.destination);
            pending_gets_.erase(found);
         }
         break;
      }
      case kError:
         handle_error(header, payload);
         break;
      case kEcho:
         break;
      case kServerDisconnect: {
         Channel *channel = find_channel_by_cid(header.parameter1);
         if (channel) {
            channel->state = ChannelState::Disconnected;
            channel->circuit = nullptr;
            channel->sid = 0;
         }
         break;
      }
      default:
         break;
      }
   }

   void handle_error(const Header &header, const uint8_t *payload)
   {
      if (header.payload_size < kHeaderSize)
         return;
      Header request;
      if (!decode_header(payload, header.payload_size, request))
         return;
      if (request.command == kReadNotify)
         pending_gets_.erase(request.parameter2);
   }

   void close_circuit(Circuit &circuit)
   {
      if (circuit.fd >= 0)
         close(circuit.fd);
      circuit.fd = -1;
      circuit.state = CircuitState::Dead;
      circuit.input.clear();
      circuit.output.clear();
      circuit.output_offset = 0;
   }

   void disconnect_circuit(Circuit &circuit)
   {
      close_circuit(circuit);
      for (Channel *channel : circuit.channels) {
         channel->circuit = nullptr;
         channel->sid = 0;
         if (channel->state == ChannelState::Creating) {
            channel->state = ChannelState::Searching;
            channel->next_search = Clock::now();
         } else {
            channel->state = ChannelState::Disconnected;
         }
      }
   }

   bool initialized_ = false;
   int udp_fd_ = -1;
   uint16_t server_port_ = kDefaultServerPort;
   uint32_t next_id_ = 1;
   uint32_t next_sequence_ = 1;
   std::vector<sockaddr_in> destinations_;
   std::vector<std::unique_ptr<Channel>> channels_;
   std::unordered_map<chid, Channel *> channel_index_;
   std::vector<std::unique_ptr<Circuit>> circuits_;
   std::unordered_map<uint32_t, PendingGet> pending_gets_;
};

ClientContext context;
std::mutex context_mutex;

int ensure_initialized()
{
   return context.initialize();
}

template <typename Operation>
int invoke(Operation operation)
{
   try {
      std::lock_guard<std::mutex> lock(context_mutex);
      int status = ensure_initialized();
      if (status != ECA_NORMAL)
         return status;
      return operation();
   } catch (const std::bad_alloc &) {
      return ECA_ALLOCMEM;
   } catch (...) {
      return ECA_INTERNAL;
   }
}

} // namespace

extern "C" int ca_task_initialize(void)
{
   return invoke([] { return ECA_NORMAL; });
}

extern "C" int ca_task_exit(void)
{
   try {
      std::lock_guard<std::mutex> lock(context_mutex);
      context.reset();
      return ECA_NORMAL;
   } catch (...) {
      return ECA_INTERNAL;
   }
}

extern "C" int ca_create_channel(const char *name, caCh *connection_handler,
                                  void *, capri priority, chid *channel)
{
   return invoke([&] {
      return context.create_channel(name, connection_handler, priority, channel);
   });
}

extern "C" int ca_pend_io(ca_real timeout)
{
   return invoke([&] { return context.pend_io(timeout); });
}

extern "C" int ca_get(chtype type, chid channel, void *value)
{
   return invoke([&] { return context.get(type, channel, value); });
}

extern "C" int ca_put(chtype type, chid channel, const void *value)
{
   return invoke([&] { return context.put(type, channel, value); });
}

extern "C" const char *ca_message(long status)
{
   switch (status) {
   case ECA_NORMAL:       return "Normal successful completion";
   case ECA_SOCK:         return "Unable to create or use a socket";
   case ECA_ALLOCMEM:     return "Unable to allocate memory";
   case ECA_TIMEOUT:      return "User specified timeout expired";
   case ECA_NOSUPPORT:    return "Feature not supported by simplified CA client";
   case ECA_BADTYPE:      return "Invalid DBR type";
   case ECA_INTERNAL:     return "Channel Access internal failure";
   case ECA_GETFAIL:      return "Channel read request failed";
   case ECA_PUTFAIL:      return "Channel write request failed";
   case ECA_BADCOUNT:     return "Invalid element count or string size";
   case ECA_BADSTR:       return "Invalid channel name";
   case ECA_DISCONN:      return "Channel is disconnected";
   case ECA_NORDACCESS:   return "Read access denied";
   case ECA_NOWTACCESS:   return "Write access denied";
   case ECA_NOSEARCHADDR: return "Empty or invalid CA search address list";
   case ECA_BADCHID:      return "Invalid channel identifier";
   case ECA_BADPRIORITY:  return "Invalid channel priority";
   default:               return "Unknown Channel Access status";
   }
}

extern "C" void ca_signal_with_file_and_lineno(long status,
                                                const char *message,
                                                const char *file, int line)
{
   std::fprintf(stderr, "CA error: %s: %s (%s:%d)\n",
                message ? message : "Channel Access operation",
                ca_message(status), file ? file : "unknown", line);
}
