Files
tcpchat/src/client.c
T
2026-02-13 21:36:45 +07:00

615 lines
15 KiB
C

#include <arpa/inet.h>
#include <getopt.h>
#include <openssl/evp.h>
#include <openssl/rand.h>
#include <signal.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <sys/ioctl.h>
#include <sys/socket.h>
#include <sys/wait.h>
#include <termios.h>
#include <unistd.h>
#define DEFAULT_PORT 7777
#define DEFAULT_IP "127.0.0.1"
#define BUFFER_SIZE 4096
#define USERNAME_MAX 50
#define INPUT_HEIGHT 3
#define DEFAULT_MSG_BUFFER 100
#define AES_KEY_SIZE 32
#define AES_IV_SIZE 12
#define AES_TAG_SIZE 16
int sock;
char username[USERNAME_MAX];
int term_rows, term_cols;
struct termios orig_termios;
int msg_buffer_size = DEFAULT_MSG_BUFFER;
char **message_history = NULL;
int history_count = 0;
int history_start = 0;
int scroll_offset = 0;
unsigned char aes_key[AES_KEY_SIZE];
// Параметры подключения
char server_ip[256] = DEFAULT_IP;
int server_port = DEFAULT_PORT;
volatile int need_redraw = 0;
volatile int running = 1;
void get_terminal_size() {
struct winsize ws;
ioctl(STDOUT_FILENO, TIOCGWINSZ, &ws);
term_rows = ws.ws_row;
term_cols = ws.ws_col;
}
void enable_raw_mode() {
tcgetattr(STDIN_FILENO, &orig_termios);
struct termios raw = orig_termios;
raw.c_lflag &= ~(ECHO | ICANON);
raw.c_cc[VMIN] = 0;
raw.c_cc[VTIME] = 1;
tcsetattr(STDIN_FILENO, TCSAFLUSH, &raw);
}
void disable_raw_mode() {
tcsetattr(STDIN_FILENO, TCSAFLUSH, &orig_termios);
}
void clear_screen() {
printf("\033[2J\033[H");
fflush(stdout);
}
void move_cursor(int row, int col) {
printf("\033[%d;%dH", row, col);
fflush(stdout);
}
void clear_line() {
printf("\033[2K");
fflush(stdout);
}
void hide_cursor() {
printf("\033[?25l");
fflush(stdout);
}
void show_cursor() {
printf("\033[?25h");
fflush(stdout);
}
void derive_key_from_password(const char *password, unsigned char *key) {
unsigned char salt[16] = {0};
if (!PKCS5_PBKDF2_HMAC(password, strlen(password), salt, sizeof(salt), 100000,
EVP_sha256(), AES_KEY_SIZE, key)) {
fprintf(stderr, "Failed to derive key\n");
exit(1);
}
}
int encrypt_message(const unsigned char *plaintext, int plaintext_len,
const unsigned char *key, const unsigned char *iv,
unsigned char *ciphertext, unsigned char *tag) {
EVP_CIPHER_CTX *ctx = EVP_CIPHER_CTX_new();
int len;
int ciphertext_len;
if (!ctx) return -1;
if (EVP_EncryptInit_ex(ctx, EVP_aes_256_gcm(), NULL, NULL, NULL) != 1)
return -1;
if (EVP_CIPHER_CTX_ctrl(ctx, EVP_CTRL_GCM_SET_IVLEN, AES_IV_SIZE, NULL) != 1)
return -1;
if (EVP_EncryptInit_ex(ctx, NULL, NULL, key, iv) != 1) return -1;
if (EVP_EncryptUpdate(ctx, ciphertext, &len, plaintext, plaintext_len) != 1)
return -1;
ciphertext_len = len;
if (EVP_EncryptFinal_ex(ctx, ciphertext + len, &len) != 1) return -1;
ciphertext_len += len;
if (EVP_CIPHER_CTX_ctrl(ctx, EVP_CTRL_GCM_GET_TAG, AES_TAG_SIZE, tag) != 1)
return -1;
EVP_CIPHER_CTX_free(ctx);
return ciphertext_len;
}
int decrypt_message(const unsigned char *ciphertext, int ciphertext_len,
const unsigned char *key, const unsigned char *iv,
const unsigned char *tag, unsigned char *plaintext) {
EVP_CIPHER_CTX *ctx = EVP_CIPHER_CTX_new();
int len;
int plaintext_len;
int ret;
if (!ctx) return -1;
if (!EVP_DecryptInit_ex(ctx, EVP_aes_256_gcm(), NULL, NULL, NULL)) {
EVP_CIPHER_CTX_free(ctx);
return -1;
}
if (!EVP_CIPHER_CTX_ctrl(ctx, EVP_CTRL_GCM_SET_IVLEN, AES_IV_SIZE, NULL)) {
EVP_CIPHER_CTX_free(ctx);
return -1;
}
if (!EVP_DecryptInit_ex(ctx, NULL, NULL, key, iv)) {
EVP_CIPHER_CTX_free(ctx);
return -1;
}
if (!EVP_DecryptUpdate(ctx, plaintext, &len, ciphertext, ciphertext_len)) {
EVP_CIPHER_CTX_free(ctx);
return -1;
}
plaintext_len = len;
if (!EVP_CIPHER_CTX_ctrl(ctx, EVP_CTRL_GCM_SET_TAG, AES_TAG_SIZE,
(void *) tag)) {
EVP_CIPHER_CTX_free(ctx);
return -1;
}
ret = EVP_DecryptFinal_ex(ctx, plaintext + len, &len);
EVP_CIPHER_CTX_free(ctx);
if (ret > 0) {
plaintext_len += len;
return plaintext_len;
} else {
return -1;
}
}
void init_message_buffer() {
message_history = calloc(msg_buffer_size, sizeof(char *));
}
void free_message_buffer() {
if (!message_history) return;
for (int i = 0; i < msg_buffer_size; i++) {
free(message_history[i]);
}
free(message_history);
}
void add_to_history(const char *msg) {
int idx = (history_start + history_count) % msg_buffer_size;
free(message_history[idx]);
message_history[idx] = strdup(msg);
if (history_count < msg_buffer_size) {
history_count++;
} else {
history_start = (history_start + 1) % msg_buffer_size;
}
}
char *get_message_by_index(int idx) {
if (idx < 0 || idx >= history_count) return NULL;
int real_idx = (history_start + idx) % msg_buffer_size;
return message_history[real_idx];
}
int get_last_visible_index() {
return history_count - 1 - scroll_offset;
}
void draw_input_frame() {
int input_row = term_rows - INPUT_HEIGHT + 1;
move_cursor(input_row, 1);
printf("+");
for (int i = 0; i < term_cols - 2; i++) printf("-");
printf("+");
move_cursor(input_row + 1, 1);
printf("| %s: ", username);
int username_len = strlen(username) + 4;
for (int i = username_len; i < term_cols - 1; i++) printf(" ");
printf("|");
move_cursor(input_row + 2, 1);
printf("+");
for (int i = 0; i < term_cols - 2; i++) printf("-");
printf("+");
if (scroll_offset > 0) {
move_cursor(input_row, term_cols - 10);
printf("[^%d]", scroll_offset);
}
fflush(stdout);
}
void redraw_chat() {
hide_cursor();
for (int i = 1; i <= term_rows - INPUT_HEIGHT; i++) {
move_cursor(i, 1);
clear_line();
}
int max_visible = term_rows - INPUT_HEIGHT;
int last_idx = get_last_visible_index();
int row = max_visible;
for (int i = 0; i < max_visible && (last_idx - i) >= 0; i++) {
char *msg = get_message_by_index(last_idx - i);
if (!msg) break;
move_cursor(row, 1);
int msg_len = strlen(msg);
if (msg_len > 0 && msg[msg_len - 1] == '\n') msg_len--;
int max_line = term_cols - 1;
if (msg_len > max_line) msg_len = max_line;
printf("%.*s", msg_len, msg);
row--;
if (row < 1) break;
}
draw_input_frame();
show_cursor();
}
void update_input_text(const char *text) {
int input_row = term_rows - INPUT_HEIGHT + 2;
move_cursor(input_row, strlen(username) + 5);
printf("%s", text);
int text_len = strlen(text);
int max_len = term_cols - strlen(username) - 6;
for (int i = text_len; i < max_len; i++) printf(" ");
move_cursor(input_row, strlen(username) + 5 + text_len);
fflush(stdout);
}
void add_chat_message(const char *msg) {
int was_at_bottom = (scroll_offset == 0);
add_to_history(msg);
if (was_at_bottom) {
scroll_offset = 0;
}
redraw_chat();
}
void scroll_up() {
int max_visible = term_rows - INPUT_HEIGHT;
int max_offset =
history_count > max_visible ? history_count - max_visible : 0;
if (scroll_offset < max_offset) {
scroll_offset++;
redraw_chat();
}
}
void scroll_down() {
if (scroll_offset > 0) {
scroll_offset--;
redraw_chat();
}
}
void send_encrypted_message(int socket, const char *plaintext) {
unsigned char iv[AES_IV_SIZE];
RAND_bytes(iv, AES_IV_SIZE);
unsigned char ciphertext[BUFFER_SIZE];
unsigned char tag[AES_TAG_SIZE];
int ciphertext_len =
encrypt_message((unsigned char *) plaintext, strlen(plaintext), aes_key,
iv, ciphertext, tag);
if (ciphertext_len < 0) {
fprintf(stderr, "Encryption failed\n");
return;
}
int total_len = AES_IV_SIZE + AES_TAG_SIZE + ciphertext_len;
int net_total_len = htonl(total_len);
send(socket, &net_total_len, sizeof(int), 0);
send(socket, iv, AES_IV_SIZE, 0);
send(socket, tag, AES_TAG_SIZE, 0);
send(socket, ciphertext, ciphertext_len, 0);
}
int receive_encrypted_message(int socket, char *plaintext, int max_len) {
int net_total_len;
int valread = read(socket, &net_total_len, sizeof(int));
if (valread <= 0) return -1;
int total_len = ntohl(net_total_len);
if (total_len <= AES_IV_SIZE + AES_TAG_SIZE || total_len > BUFFER_SIZE) {
fprintf(stderr, "Invalid total length: %d\n", total_len);
return -1;
}
unsigned char iv[AES_IV_SIZE];
unsigned char tag[AES_TAG_SIZE];
unsigned char ciphertext[BUFFER_SIZE];
valread = read(socket, iv, AES_IV_SIZE);
if (valread != AES_IV_SIZE) {
fprintf(stderr, "Failed to read IV: %d\n", valread);
return -1;
}
valread = read(socket, tag, AES_TAG_SIZE);
if (valread != AES_TAG_SIZE) {
fprintf(stderr, "Failed to read tag: %d\n", valread);
return -1;
}
int ciphertext_len = total_len - AES_IV_SIZE - AES_TAG_SIZE;
int total_read = 0;
while (total_read < ciphertext_len) {
valread =
read(socket, ciphertext + total_read, ciphertext_len - total_read);
if (valread <= 0) {
fprintf(stderr, "Failed to read ciphertext\n");
return -1;
}
total_read += valread;
}
int plaintext_len = decrypt_message(ciphertext, ciphertext_len, aes_key, iv,
tag, (unsigned char *) plaintext);
if (plaintext_len > 0) {
plaintext[plaintext_len] = '\0';
} else {
fprintf(stderr, "Decryption failed (ciphertext_len=%d)\n", ciphertext_len);
}
return plaintext_len;
}
void load_history_from_server() {
int net_count;
if (read(sock, &net_count, sizeof(int)) <= 0) return;
int count = ntohl(net_count);
printf("[INFO] Loading %d encrypted messages from history\n", count);
for (int i = 0; i < count && i < msg_buffer_size; i++) {
char plaintext[BUFFER_SIZE];
int len = receive_encrypted_message(sock, plaintext, BUFFER_SIZE);
if (len > 0) {
add_to_history(plaintext);
} else {
fprintf(stderr, "[DEBUG] Failed to decrypt history message %d\n", i);
}
}
}
void handle_input() {
char input_buffer[BUFFER_SIZE] = {0};
int pos = 0;
char c;
redraw_chat();
update_input_text("");
while (running) {
if (read(STDIN_FILENO, &c, 1) != 1) {
if (need_redraw) {
need_redraw = 0;
get_terminal_size();
redraw_chat();
update_input_text(input_buffer);
}
continue;
}
if (c == '\n') {
if (pos == 0) continue;
input_buffer[pos] = '\0';
if (strcmp(input_buffer, "/exit") == 0) {
printf("[INFO] Client is exiting\n");
running = 0;
return;
}
// Проверяем длину чтобы избежать truncation
int max_msg_len = BUFFER_SIZE - strlen(username) - 5;
if (max_msg_len < 0) max_msg_len = 0;
if ((int) strlen(input_buffer) > max_msg_len) {
input_buffer[max_msg_len] = '\0';
}
char full_msg[BUFFER_SIZE];
int written = snprintf(full_msg, sizeof(full_msg), "[%s]: %s", username,
input_buffer);
// Если truncation произошёл, ensure null-termination
if (written >= (int) sizeof(full_msg)) {
full_msg[sizeof(full_msg) - 1] = '\0';
}
send_encrypted_message(sock, full_msg);
pos = 0;
input_buffer[0] = '\0';
update_input_text("");
} else if (c == 127 || c == '\b') {
if (pos > 0) {
pos--;
input_buffer[pos] = '\0';
update_input_text(input_buffer);
}
} else if (c == 27) {
char seq[2];
if (read(STDIN_FILENO, &seq[0], 1) != 1) continue;
if (read(STDIN_FILENO, &seq[1], 1) != 1) continue;
if (seq[0] == '[') {
if (seq[1] == 'A') {
scroll_up();
} else if (seq[1] == 'B') {
scroll_down();
} else if (seq[1] == '3') {
read(STDIN_FILENO, &c, 1);
}
}
} else if (c == 21) {
pos = 0;
input_buffer[0] = '\0';
update_input_text("");
} else if (pos < BUFFER_SIZE - 1 && c >= 32 && c < 127) {
input_buffer[pos++] = c;
input_buffer[pos] = '\0';
update_input_text(input_buffer);
}
}
}
void receive_messages() {
while (running) {
char plaintext[BUFFER_SIZE];
int len = receive_encrypted_message(sock, plaintext, BUFFER_SIZE);
if (len > 0) {
add_chat_message(plaintext);
} else if (len < 0) {
add_chat_message("*** Ошибка расшифровки сообщения ***");
} else {
add_chat_message("*** Соединение закрыто сервером ***");
break;
}
}
}
void sigwinch_handler(int sig) {
need_redraw = 1;
}
void parse_args(int argc, char *argv[]) {
static struct option long_options[] = {
{"messagebuffer", required_argument, 0, 'm'},
{"ip", required_argument, 0, 'i'},
{"port", required_argument, 0, 'p'},
{0, 0, 0, 0}};
int opt;
while ((opt = getopt_long(argc, argv, "m:i:p:", long_options, NULL)) != -1) {
switch (opt) {
case 'm':
msg_buffer_size = atoi(optarg);
if (msg_buffer_size < 10) msg_buffer_size = 10;
if (msg_buffer_size > 10000) msg_buffer_size = 10000;
break;
case 'i':
strncpy(server_ip, optarg, sizeof(server_ip) - 1);
server_ip[sizeof(server_ip) - 1] = '\0';
break;
case 'p':
server_port = atoi(optarg);
if (server_port < 1 || server_port > 65535) {
fprintf(stderr, "Invalid port: %s\n", optarg);
exit(1);
}
break;
default:
fprintf(stderr,
"Usage: %s [-m N|--messagebuffer=N] [-i IP|--ip=IP] [-p "
"PORT|--port=PORT]\n",
argv[0]);
exit(1);
}
}
}
int main(int argc, char *argv[]) {
struct sockaddr_in server_addr;
parse_args(argc, argv);
printf("Enter encryption password: ");
fflush(stdout);
char password[256];
fgets(password, sizeof(password), stdin);
password[strcspn(password, "\n")] = '\0';
derive_key_from_password(password, aes_key);
printf("Enter your username: ");
fflush(stdout);
fgets(username, USERNAME_MAX, stdin);
username[strcspn(username, "\n")] = '\0';
sock = socket(AF_INET, SOCK_STREAM, 0);
if (sock < 0) {
perror("socket");
return 1;
}
server_addr.sin_family = AF_INET;
server_addr.sin_port = htons(server_port);
if (inet_pton(AF_INET, server_ip, &server_addr.sin_addr) <= 0) {
fprintf(stderr, "Invalid address: %s\n", server_ip);
return 1;
}
printf("[INFO] Connecting to %s:%d\n", server_ip, server_port);
if (connect(sock, (struct sockaddr *) &server_addr, sizeof(server_addr)) <
0) {
perror("connect");
return 1;
}
send(sock, username, USERNAME_MAX, 0);
printf("[INFO] username = %s, buffer_size = %d\n", username, msg_buffer_size);
init_message_buffer();
load_history_from_server();
get_terminal_size();
enable_raw_mode();
clear_screen();
hide_cursor();
signal(SIGWINCH, sigwinch_handler);
pid_t child_pid = fork();
if (child_pid == 0) {
receive_messages();
exit(0);
} else {
handle_input();
kill(child_pid, SIGKILL);
wait(NULL);
}
close(sock);
show_cursor();
disable_raw_mode();
clear_screen();
printf("Goodbye, %s!\n", username);
free_message_buffer();
return 0;
}