diff --git a/README.md b/README.md index 0d815eba..b1fe07ef 100644 --- a/README.md +++ b/README.md @@ -31,13 +31,18 @@ libpeer is a WebRTC implementation written in C, developed with BSD socket. The - Copy URL from the test [website](https://sepfy.github.io/libpeer) - Build and run the example ```bash -$ sudo apt -y install git cmake +$ sudo apt -y install git cmake wget ffmpeg $ git clone --recursive https://github.com/sepfy/libpeer $ cd libpeer $ cmake -S . -B build && cmake --build build -$ wget http://www.live555.com/liveMedia/public/264/test.264 # Download test video file -$ wget https://mauvecloud.net/sounds/alaw08m.wav # Download test audio file -$ ./examples/generic/sample -u +$ wget -O sample.mp4 \ + https://download.samplelib.com/mp4/sample-30s.mp4 +$ ffmpeg -i sample.mp4 \ + -map 0:v:0 -vf fps=25 -c:v libx264 -profile:v baseline -pix_fmt yuv420p \ + -x264-params bframes=0:keyint=25:min-keyint=25:scenecut=0:repeat-headers=1 \ + -f h264 test.264 \ + -map 0:a:0 -ac 1 -ar 8000 -c:a pcm_alaw -f wav test.wav +$ ./build/examples/generic/sample -u ``` - Click Connect button on the website diff --git a/examples/esp32/main/app_main.c b/examples/esp32/main/app_main.c index 4e7c0dac..86a892fa 100644 --- a/examples/esp32/main/app_main.c +++ b/examples/esp32/main/app_main.c @@ -45,7 +45,7 @@ static void oniceconnectionstatechange(PeerConnectionState state, void* user_dat ESP_LOGI(TAG, "PeerConnectionState: %d", state); eState = state; // not support datachannel close event - if (eState != PEER_CONNECTION_COMPLETED) { + if (eState != PEER_CONNECTION_CONNECTED) { gDataChannelOpened = 0; } } diff --git a/examples/esp32/main/audio.c b/examples/esp32/main/audio.c index bea78a76..1f400e06 100644 --- a/examples/esp32/main/audio.c +++ b/examples/esp32/main/audio.c @@ -122,7 +122,7 @@ void audio_task(void* arg) { ESP_LOGI(TAG, "audio task started"); for (;;) { - if (eState == PEER_CONNECTION_COMPLETED) { + if (eState == PEER_CONNECTION_CONNECTED) { ret = audio_get_samples(aenc_in_frame.buffer, aenc_in_frame.len); if (ret == aenc_in_frame.len) { diff --git a/examples/esp32/main/camera.c b/examples/esp32/main/camera.c index d961bab3..ee234a43 100644 --- a/examples/esp32/main/camera.c +++ b/examples/esp32/main/camera.c @@ -138,7 +138,7 @@ void camera_task(void* pvParameters) { last_time = get_timestamp(); for (;;) { - if ((eState == PEER_CONNECTION_COMPLETED) && gDataChannelOpened) { + if ((eState == PEER_CONNECTION_CONNECTED) && gDataChannelOpened) { fb = esp_camera_fb_get(); if (!fb) { diff --git a/examples/generic/main.c b/examples/generic/main.c index 8fc13ffd..3f3e1d7a 100644 --- a/examples/generic/main.c +++ b/examples/generic/main.c @@ -125,7 +125,7 @@ int main(int argc, char* argv[]) { reader_init(); while (!g_interrupted) { - if (g_state == PEER_CONNECTION_COMPLETED) { + if (g_state == PEER_CONNECTION_CONNECTED) { curr_time = get_timestamp(); // FPS 25 diff --git a/examples/generic/reader.c b/examples/generic/reader.c index 85ef8a57..ca6cb604 100644 --- a/examples/generic/reader.c +++ b/examples/generic/reader.c @@ -24,7 +24,7 @@ int reader_init() { FILE* video_fp = NULL; FILE* audio_fp = NULL; char videofile[] = "test.264"; - char audiofile[] = "alaw08m.wav"; + char audiofile[] = "test.wav"; video_fp = fopen(videofile, "rb"); diff --git a/examples/pico/main.c b/examples/pico/main.c index 299d64c0..9e2c0fe2 100644 --- a/examples/pico/main.c +++ b/examples/pico/main.c @@ -77,7 +77,7 @@ static void dma_i2s_in_handler(void) { } #endif - if (eState == PEER_CONNECTION_COMPLETED) { + if (eState == PEER_CONNECTION_CONNECTED) { peer_connection_send_audio(g_pc, alaw, AUDIO_BUFFER_FRAMES); } diff --git a/examples/raspberrypi/main.c b/examples/raspberrypi/main.c index 2213c14c..6aab322c 100644 --- a/examples/raspberrypi/main.c +++ b/examples/raspberrypi/main.c @@ -44,7 +44,7 @@ Media g_media; static void onconnectionstatechange(PeerConnectionState state, void* data) { printf("state is changed: %d\n", state); g_state = state; - if (g_state == PEER_CONNECTION_COMPLETED) { + if (g_state == PEER_CONNECTION_CONNECTED) { gst_element_set_state(g_media.camera_pipeline, GST_STATE_PLAYING); gst_element_set_state(g_media.mic_pipeline, GST_STATE_PLAYING); gst_element_set_state(g_media.spk_pipeline, GST_STATE_PLAYING); diff --git a/src/address.h b/src/address.h index 6f8eb781..1bd53039 100644 --- a/src/address.h +++ b/src/address.h @@ -1,14 +1,22 @@ #ifndef ADDRESS_H_ #define ADDRESS_H_ +#include + #include "config.h" -#if CONFIG_USE_LWIP + +#if CONFIG_USE_ZEPHYR +#include +#include +#include +#include +#elif CONFIG_USE_LWIP #include #else #include #include +#include #endif -#include #define ADDRSTRLEN INET6_ADDRSTRLEN diff --git a/src/agent.c b/src/agent.c index 57de5c95..d3a41a2d 100644 --- a/src/agent.c +++ b/src/agent.c @@ -63,9 +63,10 @@ static int agent_socket_recv(Agent* agent, Address* addr, uint8_t* buf, int len) int maxfd = -1; fd_set rfds; struct timeval tv; - int addr_type[] = { AF_INET, + int addr_type[] = { + AF_INET, #if CONFIG_IPV6 - AF_INET6, + AF_INET6, #endif }; @@ -126,9 +127,10 @@ static int agent_create_host_addr(Agent* agent) { int i, j; const char* iface_prefx[] = {CONFIG_IFACE_PREFIX}; IceCandidate* ice_candidate; - int addr_type[] = { AF_INET, + int addr_type[] = { + AF_INET, #if CONFIG_IPV6 - AF_INET6, + AF_INET6, #endif }; @@ -173,8 +175,8 @@ static int agent_create_stun_addr(Agent* agent, Address* serv_addr) { stun_parse_msg_buf(&recv_msg); memcpy(&bind_addr, &recv_msg.mapped_addr, sizeof(Address)); - IceCandidate* ice_candidate = agent->local_candidates + agent->local_candidates_count++; - ice_candidate_create(ice_candidate, agent->local_candidates_count, ICE_CANDIDATE_TYPE_SRFLX, &bind_addr); + IceCandidate* ice_candidate = agent->local_candidates + agent->local_candidates_count; + ice_candidate_create(ice_candidate, agent->local_candidates_count++, ICE_CANDIDATE_TYPE_SRFLX, &bind_addr); return ret; } @@ -257,6 +259,7 @@ void agent_gather_candidate(Agent* agent, const char* urls, const char* username } port = atoi(pos + 1); + printf("port => %s\n", pos + 1); if (port <= 0) { LOGE("Cannot parse port"); return; @@ -339,6 +342,26 @@ static void agent_create_binding_request(Agent* agent, StunMessage* msg) { stun_msg_finish(msg, STUN_CREDENTIAL_SHORT_TERM, agent->remote_upwd, strlen(agent->remote_upwd)); } +int agent_send_binding_request(Agent* agent) { + StunMessage msg; + StunHeader* header; + int ret; + + if (agent->nominated_pair == NULL) { + return -1; + } + + memset(&msg, 0, sizeof(msg)); + agent_create_binding_request(agent, &msg); + agent->binding_request_sent_time = ports_get_epoch_time(); + header = (StunHeader*)msg.buf; + memcpy(agent->binding_request_transaction_id, header->transaction_id, + sizeof(agent->binding_request_transaction_id)); + agent->binding_request_pending = 1; + ret = agent_socket_send(agent, &agent->nominated_pair->remote->addr, msg.buf, msg.size); + return ret; +} + void agent_process_stun_request(Agent* agent, StunMessage* stun_msg, Address* addr) { StunMessage msg; StunHeader* header; @@ -349,7 +372,6 @@ void agent_process_stun_request(Agent* agent, StunMessage* stun_msg, Address* ad memcpy(agent->transaction_id, header->transaction_id, sizeof(header->transaction_id)); agent_create_binding_response(agent, &msg, addr); agent_socket_send(agent, addr, msg.buf, msg.size); - agent->binding_request_time = ports_get_epoch_time(); } break; default: @@ -360,8 +382,13 @@ void agent_process_stun_request(Agent* agent, StunMessage* stun_msg, Address* ad void agent_process_stun_response(Agent* agent, StunMessage* stun_msg) { switch (stun_msg->stunmethod) { case STUN_METHOD_BINDING: - if (stun_msg_is_valid(stun_msg->buf, stun_msg->size, agent->remote_upwd) == 0) { + if (stun_msg_is_valid(stun_msg->buf, stun_msg->size, agent->remote_upwd) == 0 && + agent->binding_request_pending && + memcmp(((StunHeader*)stun_msg->buf)->transaction_id, + agent->binding_request_transaction_id, + sizeof(agent->binding_request_transaction_id)) == 0) { agent->nominated_pair->state = ICE_CANDIDATE_STATE_SUCCEEDED; + agent->binding_request_pending = 0; } break; default: @@ -436,7 +463,11 @@ void agent_set_remote_description(Agent* agent, char* description) { void agent_update_candidate_pairs(Agent* agent) { int i, j; + char local_addr_string[ADDRSTRLEN]; + char remote_addr_string[ADDRSTRLEN]; + int candidate_pairs_num = agent->candidate_pairs_num; // Please set gather candidates before set remote description + agent->candidate_pairs_num = 0; for (i = 0; i < agent->local_candidates_count; i++) { for (j = 0; j < agent->remote_candidates_count; j++) { if (agent->local_candidates[i].addr.family == agent->remote_candidates[j].addr.family) { @@ -448,26 +479,34 @@ void agent_update_candidate_pairs(Agent* agent) { } } } - LOGD("candidate pairs num: %d", agent->candidate_pairs_num); + + if (candidate_pairs_num != agent->candidate_pairs_num) { + LOGI("candidate pairs num %d:", agent->candidate_pairs_num); + for (i = 0; i < agent->candidate_pairs_num; i++) { + addr_to_string(&agent->candidate_pairs[i].local->addr, local_addr_string, sizeof(local_addr_string)); + addr_to_string(&agent->candidate_pairs[i].remote->addr, remote_addr_string, sizeof(remote_addr_string)); + LOGI("[%d] %s > %s", i, local_addr_string, remote_addr_string); + } + } } int agent_connectivity_check(Agent* agent) { char addr_string[ADDRSTRLEN]; uint8_t buf[1400]; - StunMessage msg; + if (agent_select_candidate_pair(agent) < 0) { + agent_update_candidate_pairs(agent); + return -1; + } if (agent->nominated_pair->state != ICE_CANDIDATE_STATE_INPROGRESS) { LOGI("nominated pair is not in progress"); return -1; } - memset(&msg, 0, sizeof(msg)); - if (agent->nominated_pair->conncheck % AGENT_CONNCHECK_PERIOD == 0) { addr_to_string(&agent->nominated_pair->remote->addr, addr_string, sizeof(addr_string)); LOGD("send binding request to remote ip: %s, port: %d", addr_string, agent->nominated_pair->remote->addr.port); - agent_create_binding_request(agent, &msg); - agent_socket_send(agent, &agent->nominated_pair->remote->addr, msg.buf, msg.size); + agent_send_binding_request(agent); } agent_recv(agent, buf, sizeof(buf)); @@ -497,7 +536,7 @@ int agent_select_candidate_pair(Agent* agent) { agent->candidate_pairs[i].state = ICE_CANDIDATE_STATE_FAILED; } else if (agent->candidate_pairs[i].state == ICE_CANDIDATE_STATE_FAILED) { } else if (agent->candidate_pairs[i].state == ICE_CANDIDATE_STATE_SUCCEEDED) { - agent->selected_pair = &agent->candidate_pairs[i]; + // agent->selected_pair = &agent->candidate_pairs[i]; return 0; } } diff --git a/src/agent.h b/src/agent.h index 382d5a4e..6ac77ba2 100644 --- a/src/agent.h +++ b/src/agent.h @@ -24,14 +24,6 @@ #define AGENT_MAX_CANDIDATE_PAIRS 100 #endif -typedef enum AgentState { - - AGENT_STATE_GATHERING_ENDED = 0, - AGENT_STATE_GATHERING_STARTED, - AGENT_STATE_GATHERING_COMPLETED, - -} AgentState; - typedef enum AgentMode { AGENT_MODE_CONTROLLED = 0, @@ -58,8 +50,9 @@ struct Agent { Address host_addr; int b_host_addr; - uint64_t binding_request_time; - AgentState state; + uint32_t binding_request_sent_time; + uint8_t binding_request_transaction_id[12]; + int binding_request_pending; AgentMode mode; @@ -88,6 +81,8 @@ int agent_select_candidate_pair(Agent* agent); int agent_connectivity_check(Agent* agent); +int agent_send_binding_request(Agent* agent); + void agent_clear_candidates(Agent* agent); int agent_create(Agent* agent); diff --git a/src/config.h b/src/config.h index 4cacfa9e..cf903aa7 100644 --- a/src/config.h +++ b/src/config.h @@ -7,6 +7,14 @@ #define SCTP_MTU (1200) #define CONFIG_MTU (1300) +#ifndef CONFIG_USE_ZEPHYR +#ifdef __ZEPHYR__ +#define CONFIG_USE_ZEPHYR 1 +#else +#define CONFIG_USE_ZEPHYR 0 +#endif +#endif + #ifndef CONFIG_USE_LWIP #define CONFIG_USE_LWIP 0 #endif @@ -49,8 +57,17 @@ #define CONFIG_TLS_READ_TIMEOUT 3000 #endif -#ifndef CONFIG_KEEPALIVE_TIMEOUT -#define CONFIG_KEEPALIVE_TIMEOUT 10000 +#ifndef CONFIG_STUN_KEEPALIVE_INTERVAL +#define CONFIG_STUN_KEEPALIVE_INTERVAL 0 +#endif + +#ifndef CONFIG_STUN_KEEPALIVE_TIMEOUT +#define CONFIG_STUN_KEEPALIVE_TIMEOUT 15000 +#endif + +#if CONFIG_STUN_KEEPALIVE_INTERVAL > 0 && \ + CONFIG_STUN_KEEPALIVE_TIMEOUT <= CONFIG_STUN_KEEPALIVE_INTERVAL +#error "CONFIG_STUN_KEEPALIVE_TIMEOUT must be greater than CONFIG_STUN_KEEPALIVE_INTERVAL" #endif #ifndef CONFIG_AUDIO_DURATION @@ -58,7 +75,7 @@ #endif #ifndef CONFIG_MAX_NALU_SIZE -#define CONFIG_MAX_NALU_SIZE (10 * 1024) // 10KB +#define CONFIG_MAX_NALU_SIZE (100 * 1024) // 100KB #endif #define CONFIG_IPV6 0 diff --git a/src/dtls_srtp.c b/src/dtls_srtp.c index dd546169..8f2ed91e 100644 --- a/src/dtls_srtp.c +++ b/src/dtls_srtp.c @@ -9,8 +9,12 @@ #if CONFIG_MBEDTLS_DEBUG #include "mbedtls/debug.h" #endif -#include "mbedtls/sha256.h" +#include "mbedtls/md.h" #include "mbedtls/ssl.h" +#include "mbedtls/version.h" +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 +#include "psa/crypto.h" +#endif #include "ports.h" #include "socket.h" #include "utils.h" @@ -21,8 +25,6 @@ int dtls_srtp_udp_send(void* ctx, const uint8_t* buf, size_t len) { int ret = udp_socket_sendto(udp_socket, dtls_srtp->remote_addr, buf, len); - LOGD("dtls_srtp_udp_send (%d)", ret); - return ret; } @@ -36,21 +38,16 @@ int dtls_srtp_udp_recv(void* ctx, uint8_t* buf, size_t len) { ports_sleep_ms(1); } - LOGD("dtls_srtp_udp_recv (%d)", ret); - return ret; } static void dtls_srtp_x509_digest(const mbedtls_x509_crt* crt, char* buf) { int i; unsigned char digest[32]; - - mbedtls_sha256_context sha256_ctx; - mbedtls_sha256_init(&sha256_ctx); - mbedtls_sha256_starts(&sha256_ctx, 0); - mbedtls_sha256_update(&sha256_ctx, crt->raw.p, crt->raw.len); - mbedtls_sha256_finish(&sha256_ctx, (unsigned char*)digest); - mbedtls_sha256_free(&sha256_ctx); + const mbedtls_md_info_t* md_info = mbedtls_md_info_from_type(MBEDTLS_MD_SHA256); + if (md_info == NULL || mbedtls_md(md_info, crt->raw.p, crt->raw.len, digest) != 0) { + memset(digest, 0, sizeof(digest)); + } for (i = 0; i < 32; i++) { snprintf(buf, 4, "%.2X:", digest[i]); @@ -66,8 +63,137 @@ static int dtls_srtp_cert_verify(void* data, mbedtls_x509_crt* crt, int depth, u return 0; } +static int dtls_srtp_generate_keypair(DtlsSrtp* dtls_srtp) { +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 + int ret; + mbedtls_svc_key_id_t key_id = MBEDTLS_SVC_KEY_ID_INIT; + psa_key_attributes_t attr = PSA_KEY_ATTRIBUTES_INIT; + + psa_set_key_usage_flags(&attr, PSA_KEY_USAGE_SIGN_HASH | PSA_KEY_USAGE_VERIFY_HASH); +#if CONFIG_DTLS_USE_ECDSA + psa_set_key_algorithm(&attr, MBEDTLS_PK_ALG_ECDSA(PSA_ALG_SHA_256)); + psa_set_key_type(&attr, PSA_KEY_TYPE_ECC_KEY_PAIR(PSA_ECC_FAMILY_SECP_R1)); + psa_set_key_bits(&attr, 256); +#else + psa_set_key_algorithm(&attr, PSA_ALG_RSA_PKCS1V15_SIGN(PSA_ALG_SHA_256)); + psa_set_key_type(&attr, PSA_KEY_TYPE_RSA_KEY_PAIR); + psa_set_key_bits(&attr, RSA_KEY_LENGTH); +#endif + if (psa_generate_key(&attr, &key_id) != PSA_SUCCESS) { + psa_reset_key_attributes(&attr); + LOGE("psa_generate_key failed"); + return -1; + } + psa_reset_key_attributes(&attr); + + ret = mbedtls_pk_wrap_psa(&dtls_srtp->pkey, key_id); + if (ret != 0) { + psa_destroy_key(key_id); + LOGE("mbedtls_pk_wrap_psa failed -0x%.4x", (unsigned int)-ret); + return ret; + } +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 + dtls_srtp->psa_key_id = key_id; +#endif + return 0; +#else +#if CONFIG_DTLS_USE_ECDSA + int ret = mbedtls_pk_setup(&dtls_srtp->pkey, mbedtls_pk_info_from_type(MBEDTLS_PK_ECKEY)); + if (ret != 0) { + return ret; + } + return mbedtls_ecp_gen_key(MBEDTLS_ECP_DP_SECP256R1, + mbedtls_pk_ec(dtls_srtp->pkey), + mbedtls_ctr_drbg_random, + &dtls_srtp->ctr_drbg); +#else + int ret = mbedtls_pk_setup(&dtls_srtp->pkey, mbedtls_pk_info_from_type(MBEDTLS_PK_RSA)); + if (ret != 0) { + return ret; + } + return mbedtls_rsa_gen_key(mbedtls_pk_rsa(dtls_srtp->pkey), + mbedtls_ctr_drbg_random, + &dtls_srtp->ctr_drbg, + RSA_KEY_LENGTH, + 65537); +#endif +#endif +} + static int dtls_srtp_selfsign_cert(DtlsSrtp* dtls_srtp) { int ret; +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 + mbedtls_x509write_cert crt; + unsigned char* cert_buf = NULL; + unsigned char serial_raw[16]; + + cert_buf = (unsigned char*)malloc(RSA_KEY_LENGTH * 2); + if (cert_buf == NULL) { + LOGE("malloc failed"); + return -1; + } + + ret = dtls_srtp_generate_keypair(dtls_srtp); + if (ret != 0) { + free(cert_buf); + return ret; + } + + mbedtls_x509write_crt_init(&crt); + mbedtls_x509write_crt_set_subject_key(&crt, &dtls_srtp->pkey); + mbedtls_x509write_crt_set_issuer_key(&crt, &dtls_srtp->pkey); + mbedtls_x509write_crt_set_version(&crt, MBEDTLS_X509_CRT_VERSION_3); + mbedtls_x509write_crt_set_md_alg(&crt, MBEDTLS_MD_SHA256); + ret = mbedtls_x509write_crt_set_subject_name(&crt, "CN=dtls_srtp"); + if (ret != 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + LOGE("mbedtls_x509write_crt_set_subject_name failed -0x%.4x", (unsigned int)-ret); + return ret; + } + ret = mbedtls_x509write_crt_set_issuer_name(&crt, "CN=dtls_srtp"); + if (ret != 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + LOGE("mbedtls_x509write_crt_set_issuer_name failed -0x%.4x", (unsigned int)-ret); + return ret; + } + + if (psa_generate_random(serial_raw, sizeof(serial_raw)) != PSA_SUCCESS) { + memset(serial_raw, 0xA5, sizeof(serial_raw)); + } + ret = mbedtls_x509write_crt_set_serial_raw(&crt, serial_raw, sizeof(serial_raw)); + if (ret != 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + LOGE("mbedtls_x509write_crt_set_serial_raw failed -0x%.4x", (unsigned int)-ret); + return ret; + } + + ret = mbedtls_x509write_crt_set_validity(&crt, "20260101000000", "20360101000000"); + if (ret != 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + LOGE("mbedtls_x509write_crt_set_validity failed -0x%.4x", (unsigned int)-ret); + return ret; + } + + ret = mbedtls_x509write_crt_pem(&crt, cert_buf, 2 * RSA_KEY_LENGTH); + if (ret != 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + LOGE("mbedtls_x509write_crt_pem failed -0x%.4x", (unsigned int)-ret); + return ret; + } + + ret = mbedtls_x509_crt_parse(&dtls_srtp->cert, cert_buf, strlen((char*)cert_buf) + 1); + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + if (ret != 0) { + LOGE("mbedtls_x509_crt_parse failed -0x%.4x", (unsigned int)-ret); + } + return ret; +#else mbedtls_x509write_cert crt; @@ -85,15 +211,22 @@ static int dtls_srtp_selfsign_cert(DtlsSrtp* dtls_srtp) { return -1; } - mbedtls_ctr_drbg_seed(&dtls_srtp->ctr_drbg, mbedtls_entropy_func, &dtls_srtp->entropy, (const unsigned char*)pers, strlen(pers)); + ret = mbedtls_ctr_drbg_seed(&dtls_srtp->ctr_drbg, + mbedtls_entropy_func, + &dtls_srtp->entropy, + (const unsigned char*)pers, + strlen(pers)); + if (ret != 0) { + free(cert_buf); + LOGE("mbedtls_ctr_drbg_seed failed -0x%.4x", (unsigned int)-ret); + return ret; + } -#if CONFIG_DTLS_USE_ECDSA - mbedtls_pk_setup(&dtls_srtp->pkey, mbedtls_pk_info_from_type(MBEDTLS_PK_ECKEY)); - mbedtls_ecp_gen_key(MBEDTLS_ECP_DP_SECP256R1, mbedtls_pk_ec(dtls_srtp->pkey), mbedtls_ctr_drbg_random, &dtls_srtp->ctr_drbg); -#else - mbedtls_pk_setup(&dtls_srtp->pkey, mbedtls_pk_info_from_type(MBEDTLS_PK_RSA)); - mbedtls_rsa_gen_key(mbedtls_pk_rsa(dtls_srtp->pkey), mbedtls_ctr_drbg_random, &dtls_srtp->ctr_drbg, RSA_KEY_LENGTH, 65537); -#endif + ret = dtls_srtp_generate_keypair(dtls_srtp); + if (ret != 0) { + free(cert_buf); + return ret; + } mbedtls_x509write_crt_init(&crt); @@ -107,9 +240,21 @@ static int dtls_srtp_selfsign_cert(DtlsSrtp* dtls_srtp) { mbedtls_x509write_crt_set_issuer_key(&crt, &dtls_srtp->pkey); - mbedtls_x509write_crt_set_subject_name(&crt, "CN=dtls_srtp"); + ret = mbedtls_x509write_crt_set_subject_name(&crt, "CN=dtls_srtp"); + if (ret != 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + LOGE("mbedtls_x509write_crt_set_subject_name failed -0x%.4x", (unsigned int)-ret); + return ret; + } - mbedtls_x509write_crt_set_issuer_name(&crt, "CN=dtls_srtp"); + ret = mbedtls_x509write_crt_set_issuer_name(&crt, "CN=dtls_srtp"); + if (ret != 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + LOGE("mbedtls_x509write_crt_set_issuer_name failed -0x%.4x", (unsigned int)-ret); + return ret; + } #if CONFIG_MBEDTLS_2_X mbedtls_mpi_init(&serial); @@ -119,24 +264,46 @@ static int dtls_srtp_selfsign_cert(DtlsSrtp* dtls_srtp) { LOGE("mbedtls_x509write_crt_set_serial failed -0x%.4x", (unsigned int)-ret); } #else - mbedtls_x509write_crt_set_serial_raw(&crt, (unsigned char*)serial, strlen(serial)); + ret = mbedtls_x509write_crt_set_serial_raw(&crt, (unsigned char*)serial, strlen(serial)); + if (ret != 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + LOGE("mbedtls_x509write_crt_set_serial_raw failed -0x%.4x", (unsigned int)-ret); + return ret; + } #endif - mbedtls_x509write_crt_set_validity(&crt, "20180101000000", "20280101000000"); + ret = mbedtls_x509write_crt_set_validity(&crt, "20260101000000", "20360101000000"); + if (ret != 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + LOGE("mbedtls_x509write_crt_set_validity failed -0x%.4x", (unsigned int)-ret); + return ret; + } ret = mbedtls_x509write_crt_pem(&crt, cert_buf, 2 * RSA_KEY_LENGTH, mbedtls_ctr_drbg_random, &dtls_srtp->ctr_drbg); if (ret < 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); LOGE("mbedtls_x509write_crt_pem failed -0x%.4x", (unsigned int)-ret); + return ret; } - mbedtls_x509_crt_parse(&dtls_srtp->cert, cert_buf, 2 * RSA_KEY_LENGTH); + ret = mbedtls_x509_crt_parse(&dtls_srtp->cert, cert_buf, strlen((char*)cert_buf) + 1); + if (ret != 0) { + mbedtls_x509write_crt_free(&crt); + free(cert_buf); + LOGE("mbedtls_x509_crt_parse failed -0x%.4x", (unsigned int)-ret); + return ret; + } mbedtls_x509write_crt_free(&crt); free(cert_buf); return ret; +#endif } #if CONFIG_MBEDTLS_DEBUG @@ -146,6 +313,7 @@ static void dtls_srtp_debug(void* ctx, int level, const char* file, int line, co #endif int dtls_srtp_init(DtlsSrtp* dtls_srtp, DtlsSrtpRole role, void* user_data) { + int ret; static const mbedtls_ssl_srtp_profile default_profiles[] = { MBEDTLS_TLS_SRTP_AES128_CM_HMAC_SHA1_80, MBEDTLS_TLS_SRTP_AES128_CM_HMAC_SHA1_32, @@ -156,6 +324,9 @@ int dtls_srtp_init(DtlsSrtp* dtls_srtp, DtlsSrtpRole role, void* user_data) { dtls_srtp->role = role; dtls_srtp->state = DTLS_SRTP_STATE_INIT; dtls_srtp->user_data = user_data; +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 + dtls_srtp->psa_key_id = MBEDTLS_SVC_KEY_ID_INIT; +#endif dtls_srtp->udp_send = dtls_srtp_udp_send; dtls_srtp->udp_recv = dtls_srtp_udp_recv; @@ -166,59 +337,110 @@ int dtls_srtp_init(DtlsSrtp* dtls_srtp, DtlsSrtpRole role, void* user_data) { mbedtls_pk_init(&dtls_srtp->pkey); mbedtls_entropy_init(&dtls_srtp->entropy); mbedtls_ctr_drbg_init(&dtls_srtp->ctr_drbg); + dtls_srtp->initialized = 1; + +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 + if (psa_crypto_init() != PSA_SUCCESS) { + LOGE("psa_crypto_init failed"); + return -1; + } +#endif + + if (dtls_srtp->role == DTLS_SRTP_ROLE_SERVER) { + ret = mbedtls_ssl_config_defaults(&dtls_srtp->conf, + MBEDTLS_SSL_IS_SERVER, + MBEDTLS_SSL_TRANSPORT_DATAGRAM, + MBEDTLS_SSL_PRESET_DEFAULT); + if (ret != 0) { + LOGE("mbedtls_ssl_config_defaults(server) failed -0x%.4x", (unsigned int)-ret); + return -1; + } + + mbedtls_ssl_cookie_init(&dtls_srtp->cookie_ctx); +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 + ret = mbedtls_ssl_cookie_setup(&dtls_srtp->cookie_ctx); +#else + ret = mbedtls_ssl_cookie_setup(&dtls_srtp->cookie_ctx, mbedtls_ctr_drbg_random, &dtls_srtp->ctr_drbg); +#endif + if (ret != 0) { + LOGE("mbedtls_ssl_cookie_setup failed -0x%.4x", (unsigned int)-ret); + return -1; + } + + mbedtls_ssl_conf_dtls_cookies(&dtls_srtp->conf, mbedtls_ssl_cookie_write, mbedtls_ssl_cookie_check, &dtls_srtp->cookie_ctx); + + } else { + ret = mbedtls_ssl_config_defaults(&dtls_srtp->conf, + MBEDTLS_SSL_IS_CLIENT, + MBEDTLS_SSL_TRANSPORT_DATAGRAM, + MBEDTLS_SSL_PRESET_DEFAULT); + if (ret != 0) { + LOGE("mbedtls_ssl_config_defaults(client) failed -0x%.4x", (unsigned int)-ret); + return -1; + } + } + #if CONFIG_MBEDTLS_DEBUG mbedtls_debug_set_threshold(3); mbedtls_ssl_conf_dbg(&dtls_srtp->conf, dtls_srtp_debug, NULL); #endif - dtls_srtp_selfsign_cert(dtls_srtp); - - mbedtls_ssl_conf_verify(&dtls_srtp->conf, dtls_srtp_cert_verify, NULL); - mbedtls_ssl_conf_authmode(&dtls_srtp->conf, MBEDTLS_SSL_VERIFY_REQUIRED); + ret = dtls_srtp_selfsign_cert(dtls_srtp); + if (ret != 0) { + LOGE("dtls_srtp_selfsign_cert failed -0x%.4x", (unsigned int)-ret); + return -1; + } + mbedtls_ssl_conf_verify(&dtls_srtp->conf, dtls_srtp_cert_verify, NULL); + /* + * WebRTC peers use self-signed certificates and authenticate them with the + * fingerprint carried in SDP. VERIFY_REQUIRED makes Mbed TLS 4.x require a + * TLS hostname as well, which does not exist for DTLS-SRTP peers. Request + * and retain the peer certificate here, then verify its SHA-256 fingerprint + * after the handshake below. + */ + mbedtls_ssl_conf_authmode(&dtls_srtp->conf, MBEDTLS_SSL_VERIFY_OPTIONAL); mbedtls_ssl_conf_ca_chain(&dtls_srtp->conf, &dtls_srtp->cert, NULL); - mbedtls_ssl_conf_own_cert(&dtls_srtp->conf, &dtls_srtp->cert, &dtls_srtp->pkey); + ret = mbedtls_ssl_conf_own_cert(&dtls_srtp->conf, &dtls_srtp->cert, &dtls_srtp->pkey); + if (ret != 0) { + LOGE("mbedtls_ssl_conf_own_cert failed -0x%.4x", (unsigned int)-ret); + return -1; + } +#if MBEDTLS_VERSION_NUMBER < 0x04000000 mbedtls_ssl_conf_rng(&dtls_srtp->conf, mbedtls_ctr_drbg_random, &dtls_srtp->ctr_drbg); - +#endif mbedtls_ssl_conf_read_timeout(&dtls_srtp->conf, 1000); - if (dtls_srtp->role == DTLS_SRTP_ROLE_SERVER) { - mbedtls_ssl_config_defaults(&dtls_srtp->conf, - MBEDTLS_SSL_IS_SERVER, - MBEDTLS_SSL_TRANSPORT_DATAGRAM, - MBEDTLS_SSL_PRESET_DEFAULT); - - mbedtls_ssl_cookie_init(&dtls_srtp->cookie_ctx); - - mbedtls_ssl_cookie_setup(&dtls_srtp->cookie_ctx, mbedtls_ctr_drbg_random, &dtls_srtp->ctr_drbg); - - mbedtls_ssl_conf_dtls_cookies(&dtls_srtp->conf, mbedtls_ssl_cookie_write, mbedtls_ssl_cookie_check, &dtls_srtp->cookie_ctx); - - } else { - mbedtls_ssl_config_defaults(&dtls_srtp->conf, - MBEDTLS_SSL_IS_CLIENT, - MBEDTLS_SSL_TRANSPORT_DATAGRAM, - MBEDTLS_SSL_PRESET_DEFAULT); - } - dtls_srtp_x509_digest(&dtls_srtp->cert, dtls_srtp->local_fingerprint); LOGD("local fingerprint: %s", dtls_srtp->local_fingerprint); - mbedtls_ssl_conf_dtls_srtp_protection_profiles(&dtls_srtp->conf, default_profiles); + ret = mbedtls_ssl_conf_dtls_srtp_protection_profiles(&dtls_srtp->conf, default_profiles); + if (ret != 0) { + LOGE("mbedtls_ssl_conf_dtls_srtp_protection_profiles failed -0x%.4x", (unsigned int)-ret); + return -1; + } mbedtls_ssl_conf_srtp_mki_value_supported(&dtls_srtp->conf, MBEDTLS_SSL_DTLS_SRTP_MKI_UNSUPPORTED); mbedtls_ssl_conf_cert_req_ca_list(&dtls_srtp->conf, MBEDTLS_SSL_CERT_REQ_CA_LIST_DISABLED); - mbedtls_ssl_setup(&dtls_srtp->ssl, &dtls_srtp->conf); + ret = mbedtls_ssl_setup(&dtls_srtp->ssl, &dtls_srtp->conf); + if (ret != 0) { + LOGE("mbedtls_ssl_setup failed -0x%.4x", (unsigned int)-ret); + return -1; + } return 0; } void dtls_srtp_deinit(DtlsSrtp* dtls_srtp) { + if (!dtls_srtp->initialized) { + return; + } + mbedtls_ssl_free(&dtls_srtp->ssl); mbedtls_ssl_config_free(&dtls_srtp->conf); @@ -227,6 +449,13 @@ void dtls_srtp_deinit(DtlsSrtp* dtls_srtp) { mbedtls_entropy_free(&dtls_srtp->entropy); mbedtls_ctr_drbg_free(&dtls_srtp->ctr_drbg); +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 + if (dtls_srtp->psa_key_id != MBEDTLS_SVC_KEY_ID_INIT) { + psa_destroy_key(dtls_srtp->psa_key_id); + dtls_srtp->psa_key_id = MBEDTLS_SVC_KEY_ID_INIT; + } +#endif + if (dtls_srtp->role == DTLS_SRTP_ROLE_SERVER) { mbedtls_ssl_cookie_free(&dtls_srtp->cookie_ctx); } @@ -235,6 +464,8 @@ void dtls_srtp_deinit(DtlsSrtp* dtls_srtp) { srtp_dealloc(dtls_srtp->srtp_in); srtp_dealloc(dtls_srtp->srtp_out); } + + dtls_srtp->initialized = 0; } static int dtls_srtp_key_derivation(DtlsSrtp* dtls_srtp, const unsigned char* master_secret, size_t secret_len, const unsigned char* randbytes, size_t randbytes_len, mbedtls_tls_prf_types tls_prf_type) { @@ -274,7 +505,7 @@ static int dtls_srtp_key_derivation(DtlsSrtp* dtls_srtp, const unsigned char* ma const uint8_t* server_key = client_key + SRTP_MASTER_KEY_LENGTH; const uint8_t* client_salt = server_key + SRTP_MASTER_KEY_LENGTH; const uint8_t* server_salt = client_salt + SRTP_MASTER_SALT_LENGTH; - uint8_t *local_key, *remote_key, *local_salt, *remote_salt; + const uint8_t *local_key, *remote_key, *local_salt, *remote_salt; if (dtls_srtp->role == DTLS_SRTP_ROLE_SERVER) { local_key = server_key; local_salt = server_salt; @@ -326,7 +557,6 @@ static int dtls_srtp_key_derivation(DtlsSrtp* dtls_srtp, const unsigned char* ma } LOGI("Created outbound SRTP session"); - dtls_srtp->state = DTLS_SRTP_STATE_CONNECTED; return 0; } @@ -395,9 +625,17 @@ static int dtls_srtp_handshake_server(DtlsSrtp* dtls_srtp) { while (1) { unsigned char client_ip[] = "test"; - mbedtls_ssl_session_reset(&dtls_srtp->ssl); + ret = mbedtls_ssl_session_reset(&dtls_srtp->ssl); + if (ret != 0) { + LOGE("mbedtls_ssl_session_reset failed -0x%.4x", (unsigned int)-ret); + break; + } - mbedtls_ssl_set_client_transport_id(&dtls_srtp->ssl, client_ip, sizeof(client_ip)); + ret = mbedtls_ssl_set_client_transport_id(&dtls_srtp->ssl, client_ip, sizeof(client_ip)); + if (ret != 0) { + LOGE("mbedtls_ssl_set_client_transport_id failed -0x%.4x", (unsigned int)-ret); + break; + } ret = dtls_srtp_do_handshake(dtls_srtp); @@ -432,9 +670,13 @@ static int dtls_srtp_handshake_client(DtlsSrtp* dtls_srtp) { return ret; } -int dtls_srtp_handshake(DtlsSrtp* dtls_srtp, Address* addr) { +int dtls_srtp_handshake(DtlsSrtp* dtls_srtp, Address* addr, const char* remote_fingerprint) { int ret; dtls_srtp->remote_addr = addr; + if (dtls_srtp->state != DTLS_SRTP_STATE_INIT) { + LOGE("DTLS-SRTP already initialized"); + return -1; + } if (dtls_srtp->role == DTLS_SRTP_ROLE_SERVER) { ret = dtls_srtp_handshake_server(dtls_srtp); @@ -442,13 +684,18 @@ int dtls_srtp_handshake(DtlsSrtp* dtls_srtp, Address* addr) { ret = dtls_srtp_handshake_client(dtls_srtp); } + if (ret != 0) { + return ret; + } + const mbedtls_x509_crt* remote_crt; if ((remote_crt = mbedtls_ssl_get_peer_cert(&dtls_srtp->ssl)) != NULL) { dtls_srtp_x509_digest(remote_crt, dtls_srtp->actual_remote_fingerprint); - if (strncmp(dtls_srtp->remote_fingerprint, dtls_srtp->actual_remote_fingerprint, DTLS_SRTP_FINGERPRINT_LENGTH) != 0) { + if (strncmp(remote_fingerprint, dtls_srtp->actual_remote_fingerprint, + DTLS_SRTP_FINGERPRINT_LENGTH) != 0) { LOGE("Actual and Expected Fingerprint mismatch: %s %s", - dtls_srtp->remote_fingerprint, + remote_fingerprint, dtls_srtp->actual_remote_fingerprint); return -1; } @@ -458,22 +705,14 @@ int dtls_srtp_handshake(DtlsSrtp* dtls_srtp, Address* addr) { return -1; } + dtls_srtp->state = DTLS_SRTP_STATE_CONNECTED; + mbedtls_dtls_srtp_info dtls_srtp_negotiation_result; mbedtls_ssl_get_dtls_srtp_negotiation_result(&dtls_srtp->ssl, &dtls_srtp_negotiation_result); return ret; } -void dtls_srtp_reset_session(DtlsSrtp* dtls_srtp) { - if (dtls_srtp->state == DTLS_SRTP_STATE_CONNECTED) { - srtp_dealloc(dtls_srtp->srtp_in); - srtp_dealloc(dtls_srtp->srtp_out); - mbedtls_ssl_session_reset(&dtls_srtp->ssl); - } - - dtls_srtp->state = DTLS_SRTP_STATE_INIT; -} - int dtls_srtp_write(DtlsSrtp* dtls_srtp, const unsigned char* buf, size_t len) { int ret; @@ -514,8 +753,8 @@ void dtls_srtp_decrypt_rtcp_packet(DtlsSrtp* dtls_srtp, uint8_t* packet, int* by srtp_unprotect_rtcp(dtls_srtp->srtp_in, packet, bytes); } -void dtls_srtp_encrypt_rtp_packet(DtlsSrtp* dtls_srtp, uint8_t* packet, int* bytes) { - srtp_protect(dtls_srtp->srtp_out, packet, bytes); +int dtls_srtp_encrypt_rtp_packet(DtlsSrtp* dtls_srtp, uint8_t* packet, int* bytes) { + return (int)srtp_protect(dtls_srtp->srtp_out, packet, bytes); } void dtls_srtp_encrypt_rctp_packet(DtlsSrtp* dtls_srtp, uint8_t* packet, int* bytes) { diff --git a/src/dtls_srtp.h b/src/dtls_srtp.h index 09d24017..443edf79 100644 --- a/src/dtls_srtp.h +++ b/src/dtls_srtp.h @@ -4,8 +4,21 @@ #include #include +#if defined(__has_include) +#if __has_include() #include +#else +#include +#endif +#if __has_include() #include +#else +#include +#endif +#else +#include +#include +#endif #include #include #include @@ -44,6 +57,9 @@ typedef struct DtlsSrtp { mbedtls_ssl_cookie_ctx cookie_ctx; mbedtls_x509_crt cert; mbedtls_pk_context pkey; +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 + mbedtls_svc_key_id_t psa_key_id; +#endif mbedtls_entropy_context entropy; mbedtls_ctr_drbg_context ctr_drbg; @@ -62,9 +78,9 @@ typedef struct DtlsSrtp { DtlsSrtpRole role; DtlsSrtpState state; + int initialized; char local_fingerprint[DTLS_SRTP_FINGERPRINT_LENGTH]; - char remote_fingerprint[DTLS_SRTP_FINGERPRINT_LENGTH]; char actual_remote_fingerprint[DTLS_SRTP_FINGERPRINT_LENGTH]; void* user_data; @@ -77,9 +93,7 @@ void dtls_srtp_deinit(DtlsSrtp* dtls_srtp); int dtls_srtp_create_cert(DtlsSrtp* dtls_srtp); -int dtls_srtp_handshake(DtlsSrtp* dtls_srtp, Address* addr); - -void dtls_srtp_reset_session(DtlsSrtp* dtls_srtp); +int dtls_srtp_handshake(DtlsSrtp* dtls_srtp, Address* addr, const char* remote_fingerprint); int dtls_srtp_write(DtlsSrtp* dtls_srtp, const uint8_t* buf, size_t len); @@ -93,7 +107,7 @@ void dtls_srtp_decrypt_rtp_packet(DtlsSrtp* dtls_srtp, uint8_t* packet, int* byt void dtls_srtp_decrypt_rtcp_packet(DtlsSrtp* dtls_srtp, uint8_t* packet, int* bytes); -void dtls_srtp_encrypt_rtp_packet(DtlsSrtp* dtls_srtp, uint8_t* packet, int* bytes); +int dtls_srtp_encrypt_rtp_packet(DtlsSrtp* dtls_srtp, uint8_t* packet, int* bytes); void dtls_srtp_encrypt_rctp_packet(DtlsSrtp* dtls_srtp, uint8_t* packet, int* bytes); diff --git a/src/peer_connection.c b/src/peer_connection.c index 5b1f5fb1..c4fab11d 100644 --- a/src/peer_connection.c +++ b/src/peer_connection.c @@ -25,6 +25,8 @@ struct PeerConnection { Agent agent; DtlsSrtp dtls_srtp; Sctp sctp; + DtlsSrtpRole role; + char remote_fingerprint[DTLS_SRTP_FINGERPRINT_LENGTH]; char sdp[CONFIG_SDP_BUFFER_SIZE]; @@ -36,7 +38,6 @@ struct PeerConnection { uint8_t temp_buf[CONFIG_MTU]; uint8_t agent_buf[CONFIG_MTU]; int agent_ret; - int b_local_description_created; RtpEncoder artp_encoder; RtpEncoder vrtp_encoder; @@ -49,8 +50,13 @@ struct PeerConnection { static void peer_connection_outgoing_rtp_packet(uint8_t* data, size_t size, void* user_data) { PeerConnection* pc = (PeerConnection*)user_data; - dtls_srtp_encrypt_rtp_packet(&pc->dtls_srtp, data, (int*)&size); - agent_send(&pc->agent, data, size); + int packet_len = (int)size; + + if (dtls_srtp_encrypt_rtp_packet(&pc->dtls_srtp, data, &packet_len) != 0) { + return; + } + + agent_send(&pc->agent, data, packet_len); } static int peer_connection_dtls_srtp_recv(void* ctx, unsigned char* buf, size_t len) { @@ -64,7 +70,7 @@ static int peer_connection_dtls_srtp_recv(void* ctx, unsigned char* buf, size_t return pc->agent_ret; } - while (recv_max < CONFIG_TLS_READ_TIMEOUT && pc->state == PEER_CONNECTION_CONNECTED) { + while (recv_max < CONFIG_TLS_READ_TIMEOUT && pc->state == PEER_CONNECTION_CHECKING) { ret = agent_recv(&pc->agent, buf, len); if (ret > 0) { @@ -94,7 +100,7 @@ static void peer_connection_incoming_rtcp(PeerConnection* pc, uint8_t* buf, size switch (rtcp_header->type) { case RTCP_RR: LOGD("RTCP_PR"); - if (rtcp_header->rc > 0) { + if (rtcp_header_rc(rtcp_header) > 0) { // TODO: REMB, GCC ...etc #if 0 RtcpRr rtcp_rr = rtcp_parse_rr(buf); @@ -108,7 +114,7 @@ static void peer_connection_incoming_rtcp(PeerConnection* pc, uint8_t* buf, size } break; case RTCP_PSFB: { - int fmt = rtcp_header->rc; + int fmt = rtcp_header_rc(rtcp_header); LOGD("RTCP_PSFB %d", fmt); // PLI and FIR if ((fmt == 1 || fmt == 4) && pc->config.on_request_keyframe) { @@ -131,8 +137,6 @@ const char* peer_connection_state_to_string(PeerConnectionState state) { return "checking"; case PEER_CONNECTION_CONNECTED: return "connected"; - case PEER_CONNECTION_COMPLETED: - return "completed"; case PEER_CONNECTION_FAILED: return "failed"; case PEER_CONNECTION_CLOSED: @@ -198,7 +202,7 @@ void peer_connection_close(PeerConnection* pc) { } int peer_connection_send_audio(PeerConnection* pc, const uint8_t* buf, size_t len) { - if (pc->state != PEER_CONNECTION_COMPLETED) { + if (pc->state != PEER_CONNECTION_CONNECTED) { // LOGE("dtls_srtp not connected"); return -1; } @@ -206,7 +210,7 @@ int peer_connection_send_audio(PeerConnection* pc, const uint8_t* buf, size_t le } int peer_connection_send_video(PeerConnection* pc, const uint8_t* buf, size_t len) { - if (pc->state != PEER_CONNECTION_COMPLETED) { + if (pc->state != PEER_CONNECTION_CONNECTED) { // LOGE("dtls_srtp not connected"); return -1; } @@ -259,12 +263,13 @@ int peer_connection_create_datachannel_sid(PeerConnection* pc, DecpChannelType c // +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ int msg_size = 12 + strlen(label) + strlen(protocol); uint16_t priority_big_endian = htons(priority); - uint32_t reliability_big_endian = ntohl(reliability_parameter); + uint32_t reliability_big_endian = htonl(reliability_parameter); uint16_t label_length = htons(strlen(label)); uint16_t protocol_length = htons(strlen(protocol)); char* msg = calloc(1, msg_size); msg[0] = DATA_CHANNEL_OPEN; + msg[1] = (uint8_t)channel_type; memcpy(msg + 2, &priority_big_endian, sizeof(uint16_t)); memcpy(msg + 4, &reliability_big_endian, sizeof(uint32_t)); memcpy(msg + 8, &label_length, sizeof(uint16_t)); @@ -285,34 +290,31 @@ int peer_connection_loop(PeerConnection* pc) { uint32_t ssrc = 0; memset(pc->agent_buf, 0, sizeof(pc->agent_buf)); pc->agent_ret = -1; - switch (pc->state) { case PEER_CONNECTION_NEW: break; case PEER_CONNECTION_CHECKING: - if (agent_select_candidate_pair(&pc->agent) < 0) { - STATE_CHANGED(pc, PEER_CONNECTION_FAILED); - } else if (agent_connectivity_check(&pc->agent) == 0) { - STATE_CHANGED(pc, PEER_CONNECTION_CONNECTED); + if (pc->agent.selected_pair) { + // if ice candidate pass the connectivity check, then we can start DTLS-SRTP handshake + dtls_srtp_handshake(&pc->dtls_srtp, NULL, + pc->remote_fingerprint); + } else { + agent_connectivity_check(&pc->agent); } - break; - - case PEER_CONNECTION_CONNECTED: - if (dtls_srtp_handshake(&pc->dtls_srtp, NULL) == 0) { + if (pc->dtls_srtp.state == DTLS_SRTP_STATE_CONNECTED) { LOGD("DTLS-SRTP handshake done"); - if (pc->config.datachannel) { - LOGI("SCTP create socket"); + LOGI("Creating SCTP association"); sctp_create_association(&pc->sctp, &pc->dtls_srtp); pc->sctp.userdata = pc->config.user_data; } - STATE_CHANGED(pc, PEER_CONNECTION_COMPLETED); + STATE_CHANGED(pc, PEER_CONNECTION_CONNECTED); } break; - case PEER_CONNECTION_COMPLETED: + case PEER_CONNECTION_CONNECTED: if ((pc->agent_ret = agent_recv(&pc->agent, pc->agent_buf, sizeof(pc->agent_buf))) > 0) { LOGD("agent_recv %d", pc->agent_ret); @@ -346,10 +348,25 @@ int peer_connection_loop(PeerConnection* pc) { } } - if (CONFIG_KEEPALIVE_TIMEOUT > 0 && (ports_get_epoch_time() - pc->agent.binding_request_time) > CONFIG_KEEPALIVE_TIMEOUT) { - LOGI("binding request timeout"); - STATE_CHANGED(pc, PEER_CONNECTION_CLOSED); +#if CONFIG_STUN_KEEPALIVE_INTERVAL > 0 + { + uint32_t elapsed = + (uint32_t)(ports_get_epoch_time() - pc->agent.binding_request_sent_time); + + if (pc->agent.binding_request_pending) { + if (elapsed >= CONFIG_STUN_KEEPALIVE_TIMEOUT) { + LOGW("STUN keepalive response timeout"); + STATE_CHANGED(pc, PEER_CONNECTION_CLOSED); + } + } else if (elapsed >= CONFIG_STUN_KEEPALIVE_INTERVAL) { + if (agent_send_binding_request(&pc->agent) < 0) { + LOGW("Failed to send STUN keepalive"); + } else { + LOGD("Sent STUN keepalive"); + } + } } +#endif break; case PEER_CONNECTION_FAILED: @@ -385,7 +402,14 @@ void peer_connection_set_remote_description(PeerConnection* pc, const char* sdp, } if (strstr(buf, "a=fingerprint")) { - strncpy(pc->dtls_srtp.remote_fingerprint, buf + 22, DTLS_SRTP_FINGERPRINT_LENGTH); + const char* fingerprint = buf + 22; + size_t fingerprint_len = strlen(fingerprint); + + if (fingerprint_len >= sizeof(pc->remote_fingerprint)) { + LOGE("remote fingerprint is too long"); + return; + } + memcpy(pc->remote_fingerprint, fingerprint, fingerprint_len + 1); } if (strstr(buf, "a=ice-ufrag") && @@ -403,6 +427,7 @@ void peer_connection_set_remote_description(PeerConnection* pc, const char* sdp, if ((val_start = strstr(buf, "a=ssrc:")) && ssrc) { *ssrc = strtoul(val_start + 7, NULL, 10); LOGD("SSRC: %" PRIu32, *ssrc); + ssrc = NULL; } start = line + 2; @@ -413,102 +438,116 @@ void peer_connection_set_remote_description(PeerConnection* pc, const char* sdp, } agent_set_remote_description(&pc->agent, (char*)sdp); - if (type == SDP_TYPE_ANSWER) { - agent_update_candidate_pairs(&pc->agent); + agent_update_candidate_pairs(&pc->agent); + if (pc->state == PEER_CONNECTION_NEW) { STATE_CHANGED(pc, PEER_CONNECTION_CHECKING); } } -static const char* peer_connection_create_sdp(PeerConnection* pc, SdpType sdp_type) { - char* description = (char*)pc->temp_buf; - - memset(pc->temp_buf, 0, sizeof(pc->temp_buf)); - DtlsSrtpRole role = DTLS_SRTP_ROLE_SERVER; - +void peer_connection_set_local_description(PeerConnection* pc, const char* sdp, SdpType sdp_type) { + // just for gathering ICE candidates + if (pc->state == PEER_CONNECTION_CONNECTED) { + return; + } pc->sctp.connected = 0; + dtls_srtp_deinit(&pc->dtls_srtp); + memset(&pc->dtls_srtp, 0, sizeof(pc->dtls_srtp)); + switch (sdp_type) { case SDP_TYPE_OFFER: - role = DTLS_SRTP_ROLE_SERVER; + pc->role = DTLS_SRTP_ROLE_SERVER; agent_clear_candidates(&pc->agent); pc->agent.mode = AGENT_MODE_CONTROLLING; break; case SDP_TYPE_ANSWER: - role = DTLS_SRTP_ROLE_CLIENT; + pc->role = DTLS_SRTP_ROLE_CLIENT; pc->agent.mode = AGENT_MODE_CONTROLLED; break; default: break; } - dtls_srtp_reset_session(&pc->dtls_srtp); - dtls_srtp_init(&pc->dtls_srtp, role, pc); + if (dtls_srtp_init(&pc->dtls_srtp, pc->role, pc) != 0) { + LOGE("dtls_srtp_init failed"); + STATE_CHANGED(pc, PEER_CONNECTION_FAILED); + return; + } pc->dtls_srtp.udp_recv = peer_connection_dtls_srtp_recv; pc->dtls_srtp.udp_send = peer_connection_dtls_srtp_send; + agent_create_ice_credential(&pc->agent); + agent_gather_candidate(&pc->agent, NULL, NULL, NULL); // host address + for (int i = 0; i < sizeof(pc->config.ice_servers) / sizeof(pc->config.ice_servers[0]); ++i) { + if (pc->config.ice_servers[i].urls) { + LOGI("ice server: %s", pc->config.ice_servers[i].urls); + agent_gather_candidate(&pc->agent, pc->config.ice_servers[i].urls, pc->config.ice_servers[i].username, pc->config.ice_servers[i].credential); + } + } +} + +static const char* peer_connection_create_sdp(PeerConnection* pc, SdpType sdp_type) { + char* description = (char*)pc->temp_buf; + memset(pc->temp_buf, 0, sizeof(pc->temp_buf)); memset(pc->sdp, 0, sizeof(pc->sdp)); // TODO: check if we have video or audio codecs - sdp_create(pc->sdp, - pc->config.video_codec != CODEC_NONE, - pc->config.audio_codec != CODEC_NONE, - pc->config.datachannel); + int sdp_audio = (pc->config.audio_codec != CODEC_NONE) && (sdp_type == SDP_TYPE_OFFER || pc->remote_assrc > 0); + int sdp_video = (pc->config.video_codec != CODEC_NONE) && (sdp_type == SDP_TYPE_OFFER || pc->remote_vssrc > 0); + + sdp_create(pc->sdp, sdp_video, sdp_audio, pc->config.datachannel); - agent_create_ice_credential(&pc->agent); sdp_append(pc->sdp, "a=ice-ufrag:%s", pc->agent.local_ufrag); sdp_append(pc->sdp, "a=ice-pwd:%s", pc->agent.local_upwd); sdp_append(pc->sdp, "a=fingerprint:sha-256 %s", pc->dtls_srtp.local_fingerprint); - sdp_append(pc->sdp, peer_connection_dtls_role_setup_value(role)); + sdp_append(pc->sdp, peer_connection_dtls_role_setup_value(pc->role)); - if (pc->config.video_codec == CODEC_H264) { - sdp_append_h264(pc->sdp); + if (sdp_video) { + switch (pc->config.video_codec) { + case CODEC_H264: + sdp_append_h264(pc->sdp); + break; + case CODEC_VP8: + sdp_append_vp8(pc->sdp); + break; + } } - switch (pc->config.audio_codec) { - case CODEC_PCMA: - sdp_append_pcma(pc->sdp); - break; - case CODEC_PCMU: - sdp_append_pcmu(pc->sdp); - break; - case CODEC_OPUS: - sdp_append_opus(pc->sdp); - default: - break; + if (sdp_audio) { + switch (pc->config.audio_codec) { + case CODEC_PCMA: + sdp_append_pcma(pc->sdp); + break; + case CODEC_PCMU: + sdp_append_pcmu(pc->sdp); + break; + case CODEC_OPUS: + sdp_append_opus(pc->sdp); + default: + break; + } } if (pc->config.datachannel) { sdp_append_datachannel(pc->sdp); } - pc->b_local_description_created = 1; - - agent_gather_candidate(&pc->agent, NULL, NULL, NULL); // host address - for (int i = 0; i < sizeof(pc->config.ice_servers) / sizeof(pc->config.ice_servers[0]); ++i) { - if (pc->config.ice_servers[i].urls) { - LOGI("ice server: %s", pc->config.ice_servers[i].urls); - agent_gather_candidate(&pc->agent, pc->config.ice_servers[i].urls, pc->config.ice_servers[i].username, pc->config.ice_servers[i].credential); - } - } - agent_get_local_description(&pc->agent, description, sizeof(pc->temp_buf)); sdp_append(pc->sdp, description); if (pc->onicecandidate) { pc->onicecandidate(pc->sdp, pc->config.user_data); } - return pc->sdp; } const char* peer_connection_create_offer(PeerConnection* pc) { + peer_connection_set_local_description(pc, NULL, SDP_TYPE_OFFER); return peer_connection_create_sdp(pc, SDP_TYPE_OFFER); } const char* peer_connection_create_answer(PeerConnection* pc) { - const char* sdp = peer_connection_create_sdp(pc, SDP_TYPE_ANSWER); - agent_update_candidate_pairs(&pc->agent); - STATE_CHANGED(pc, PEER_CONNECTION_CHECKING); - return sdp; + peer_connection_set_local_description(pc, NULL, SDP_TYPE_ANSWER); + return peer_connection_create_sdp(pc, SDP_TYPE_ANSWER); } int peer_connection_send_rtcp_pil(PeerConnection* pc, uint32_t ssrc) { @@ -578,6 +617,7 @@ int peer_connection_add_ice_candidate(PeerConnection* pc, char* candidate) { if (ice_candidate_from_description(&agent->remote_candidates[agent->remote_candidates_count], candidate, candidate + strlen(candidate)) != 0) { return -1; } + LOGD("Add candidate: %s", candidate); agent->remote_candidates_count++; return 0; diff --git a/src/peer_connection.h b/src/peer_connection.h index b087893e..600b6a57 100644 --- a/src/peer_connection.h +++ b/src/peer_connection.h @@ -18,15 +18,12 @@ typedef enum SdpType { } SdpType; typedef enum PeerConnectionState { - - PEER_CONNECTION_CLOSED = 0, PEER_CONNECTION_NEW, PEER_CONNECTION_CHECKING, PEER_CONNECTION_CONNECTED, - PEER_CONNECTION_COMPLETED, - PEER_CONNECTION_FAILED, PEER_CONNECTION_DISCONNECTED, - + PEER_CONNECTION_FAILED, + PEER_CONNECTION_CLOSED, } PeerConnectionState; typedef enum DataChannelType { diff --git a/src/ports.c b/src/ports.c index 2346f609..a54ccd5f 100644 --- a/src/ports.c +++ b/src/ports.c @@ -11,6 +11,9 @@ #include "lwip/netdb.h" #include "lwip/netif.h" #include "lwip/sys.h" +#elif CONFIG_USE_ZEPHYR +#include +#include #else #include #include @@ -51,6 +54,29 @@ int ports_get_host_addr(Address* addr, const char* iface_prefix) { break; } } +#elif CONFIG_USE_ZEPHYR + struct net_if* iface; + struct in_addr* ifaddr; + uint16_t port; + + ARG_UNUSED(iface_prefix); + port = addr->port; + iface = net_if_get_default(); + if (iface == NULL) { + LOGE("No default network interface"); + return 0; + } + + ifaddr = net_if_ipv4_get_global_addr(iface, NET_ADDR_PREFERRED); + if (ifaddr == NULL) { + LOGE("No global IPv4 address on default interface"); + return 0; + } + + addr_set_family(addr, AF_INET); + addr_set_port(addr, port); + addr->sin.sin_addr = *ifaddr; + ret = 1; #else struct ifaddrs *ifaddr, *ifa; diff --git a/src/rtcp.c b/src/rtcp.c index 117b58c5..66d5f9e4 100644 --- a/src/rtcp.c +++ b/src/rtcp.c @@ -10,7 +10,8 @@ int rtcp_probe(uint8_t* packet, size_t size) { return -1; RtpHeader* header = (RtpHeader*)packet; - return ((header->type >= 64) && (header->type < 96)); + uint8_t type = rtp_header_type(header); + return ((type >= 64) && (type < 96)); } int rtcp_get_pli(uint8_t* packet, int len, uint32_t ssrc) { @@ -19,9 +20,7 @@ int rtcp_get_pli(uint8_t* packet, int len, uint32_t ssrc) { memset(packet, 0, len); RtcpHeader* rtcp_header = (RtcpHeader*)packet; - rtcp_header->version = 2; - rtcp_header->type = RTCP_PSFB; - rtcp_header->rc = 1; + rtcp_header_init(rtcp_header, RTCP_PSFB, 1); rtcp_header->length = htons((len / 4) - 1); memcpy(packet + 8, &ssrc, 4); @@ -38,9 +37,7 @@ int rtcp_get_fir(uint8_t* packet, int len, int* seqnr) { if (*seqnr < 0 || *seqnr >= 256) *seqnr = 0; - rtcp->version = 2; - rtcp->type = RTCP_PSFB; - rtcp->rc = 4; + rtcp_header_init(rtcp, RTCP_PSFB, 4); rtcp->length = htons((len / 4) - 1); RtcpFb* rtcp_fb = (RtcpFb*)rtcp; RtcpFir* fir = (RtcpFir*)rtcp_fb->fci; diff --git a/src/rtcp.h b/src/rtcp.h index f888d057..0079b67c 100644 --- a/src/rtcp.h +++ b/src/rtcp.h @@ -1,14 +1,8 @@ #ifndef RTCP_H_ #define RTCP_H_ -#ifdef __BYTE_ORDER -#define __BIG_ENDIAN 4321 -#define __LITTLE_ENDIAN 1234 -#elif __APPLE__ -#include -#else -#include -#endif +#include +#include typedef enum RtcpType { @@ -25,21 +19,29 @@ typedef enum RtcpType { } RtcpType; typedef struct RtcpHeader { -#if __BYTE_ORDER == __BIG_ENDIAN - uint16_t version : 2; - uint16_t padding : 1; - uint16_t rc : 5; - uint16_t type : 8; -#elif __BYTE_ORDER == __LITTLE_ENDIAN - uint16_t rc : 5; - uint16_t padding : 1; - uint16_t version : 2; - uint16_t type : 8; -#endif - uint16_t length : 16; + uint8_t vprc; /* Version, Padding, Report count/Feedback message type */ + uint8_t type; /* Packet Type */ + uint16_t length; } RtcpHeader; +static inline void rtcp_header_init(RtcpHeader* header, uint8_t type, uint8_t rc) { + header->vprc = 0x80U | (rc & 0x1fU); + header->type = type; +} + +static inline uint8_t rtcp_header_version(const RtcpHeader* header) { + return (header->vprc >> 6) & 0x03U; +} + +static inline uint8_t rtcp_header_padding(const RtcpHeader* header) { + return (header->vprc >> 5) & 0x01U; +} + +static inline uint8_t rtcp_header_rc(const RtcpHeader* header) { + return header->vprc & 0x1fU; +} + typedef struct RtcpReportBlock { uint32_t ssrc; uint32_t flcnpl; diff --git a/src/rtp.c b/src/rtp.c index 0a3136ad..21d29f71 100644 --- a/src/rtp.c +++ b/src/rtp.c @@ -10,6 +10,7 @@ typedef enum RtpH264Type { NALU = 23, + STAP_A = 24, FU_A = 28, } RtpH264Type; @@ -29,13 +30,17 @@ typedef struct FuHeader { #define RTP_PAYLOAD_SIZE (CONFIG_MTU - sizeof(RtpHeader)) #define FU_PAYLOAD_SIZE (CONFIG_MTU - sizeof(RtpHeader) - sizeof(FuHeader) - sizeof(NaluHeader)) +#define NALU_START_CODE_SIZE 4 +#define NALU_HEADER_SIZE 1 +#define FU_HEADER_SIZE 1 int rtp_packet_validate(uint8_t* packet, size_t size) { if (size < 12) return 0; RtpHeader* rtp_header = (RtpHeader*)packet; - return ((rtp_header->type < 64) || (rtp_header->type >= 96)); + uint8_t type = rtp_header_type(rtp_header); + return ((type < 64) || (type >= 96)); } uint32_t rtp_get_ssrc(uint8_t* packet) { @@ -46,19 +51,15 @@ uint32_t rtp_get_ssrc(uint8_t* packet) { static int rtp_encoder_encode_h264_single(RtpEncoder* rtp_encoder, uint8_t* buf, size_t size) { RtpPacket* rtp_packet = (RtpPacket*)rtp_encoder->buf; - rtp_packet->header.version = 2; - rtp_packet->header.padding = 0; - rtp_packet->header.extension = 0; - rtp_packet->header.csrccount = 0; - rtp_packet->header.markerbit = 0; - rtp_packet->header.type = rtp_encoder->type; - rtp_packet->header.seq_number = htons(rtp_encoder->seq_number++); + rtp_header_init(&rtp_packet->header, rtp_encoder->type); + rtp_packet->header.seq_number = htons(rtp_encoder->seq_number); + rtp_encoder->seq_number++; rtp_packet->header.timestamp = htonl(rtp_encoder->timestamp); rtp_packet->header.ssrc = htonl(rtp_encoder->ssrc); // I frame and P frame if ((*buf & 0x1f) == 0x05 || (*buf & 0x1f) == 0x01) { - rtp_packet->header.markerbit = 1; + rtp_header_set_marker(&rtp_packet->header); rtp_encoder->timestamp += rtp_encoder->timestamp_increment; } #if 0 @@ -73,12 +74,7 @@ static int rtp_encoder_encode_h264_single(RtpEncoder* rtp_encoder, uint8_t* buf, static int rtp_encoder_encode_h264_fu_a(RtpEncoder* rtp_encoder, uint8_t* buf, size_t size) { RtpPacket* rtp_packet = (RtpPacket*)rtp_encoder->buf; - rtp_packet->header.version = 2; - rtp_packet->header.padding = 0; - rtp_packet->header.extension = 0; - rtp_packet->header.csrccount = 0; - rtp_packet->header.markerbit = 0; - rtp_packet->header.type = rtp_encoder->type; + rtp_header_init(&rtp_packet->header, rtp_encoder->type); rtp_packet->header.timestamp = htonl(rtp_encoder->timestamp); rtp_packet->header.ssrc = htonl(rtp_encoder->ssrc); uint8_t type = buf[0] & 0x1f; @@ -101,11 +97,12 @@ static int rtp_encoder_encode_h264_fu_a(RtpEncoder* rtp_encoder, uint8_t* buf, s fu_indicator->f = 0; fu_header->type = type; fu_header->r = 0; - rtp_packet->header.seq_number = htons(rtp_encoder->seq_number++); + rtp_packet->header.seq_number = htons(rtp_encoder->seq_number); + rtp_encoder->seq_number++; if (size <= FU_PAYLOAD_SIZE) { fu_header->e = 1; - rtp_packet->header.markerbit = 1; + rtp_header_set_marker(&rtp_packet->header); memcpy(rtp_packet->payload + sizeof(NaluHeader) + sizeof(FuHeader), buf, size); rtp_encoder->on_packet(rtp_encoder->buf, size + sizeof(RtpHeader) + sizeof(NaluHeader) + sizeof(FuHeader), rtp_encoder->user_data); break; @@ -163,13 +160,9 @@ static int rtp_encoder_encode_h264(RtpEncoder* rtp_encoder, uint8_t* buf, size_t static int rtp_encoder_encode_generic(RtpEncoder* rtp_encoder, uint8_t* buf, size_t size) { RtpHeader* rtp_header = (RtpHeader*)rtp_encoder->buf; - rtp_header->version = 2; - rtp_header->padding = 0; - rtp_header->extension = 0; - rtp_header->csrccount = 0; - rtp_header->markerbit = 0; - rtp_header->type = rtp_encoder->type; - rtp_header->seq_number = htons(rtp_encoder->seq_number++); + rtp_header_init(rtp_header, rtp_encoder->type); + rtp_header->seq_number = htons(rtp_encoder->seq_number); + rtp_encoder->seq_number++; rtp_header->timestamp = htonl(rtp_encoder->timestamp); rtp_encoder->timestamp += rtp_encoder->timestamp_increment; rtp_header->ssrc = htonl(rtp_encoder->ssrc); @@ -220,48 +213,98 @@ int rtp_encoder_encode(RtpEncoder* rtp_encoder, const uint8_t* buf, size_t size) return rtp_encoder->encode_func(rtp_encoder, (uint8_t*)buf, size); } +static const uint32_t nalu_start_4bytecode = 0x01000000; +static int rtp_decode_h264_stap_a(RtpDecoder* rtp_decoder, + uint8_t* buf, + size_t size, + uint8_t* nalu_buf, + int* nalu_offset) { + *nalu_offset = 0; + while (*nalu_offset + 2 < size) { + uint16_t nalu_length = buf[*nalu_offset] << 8 | buf[*nalu_offset + 1]; + *nalu_offset += 2; + + if (*nalu_offset + nalu_length > size) { + LOGE("Invalid STAP-A packet: NALU length exceeds packet size"); + return -1; + } + + memcpy(nalu_buf, &nalu_start_4bytecode, NALU_START_CODE_SIZE); + memcpy(nalu_buf + NALU_START_CODE_SIZE, buf + *nalu_offset, nalu_length); + + if (rtp_decoder->on_packet != NULL) { + rtp_decoder->on_packet(nalu_buf, NALU_START_CODE_SIZE + nalu_length, rtp_decoder->user_data); + } + + *nalu_offset += nalu_length; + } + return 0; +} + +static int rtp_decode_h264_single(RtpDecoder* rtp_decoder, + uint8_t* buf, + size_t size, + uint8_t* nalu_buf, + int* nalu_offset) { + memcpy(nalu_buf, &nalu_start_4bytecode, NALU_START_CODE_SIZE); + *nalu_offset = NALU_START_CODE_SIZE; + memcpy(nalu_buf + *nalu_offset, buf, size); + *nalu_offset += size; + if (rtp_decoder->on_packet != NULL) { + rtp_decoder->on_packet(nalu_buf, *nalu_offset, rtp_decoder->user_data); + } + *nalu_offset = 0; // reset for next NALU + return 0; +} + +static int rtp_decode_h264_fu_a(RtpDecoder* rtp_decoder, + uint8_t* buf, + size_t size, + uint8_t* nalu_buf, + int* nalu_offset) { + NaluHeader* fu_indicator = (NaluHeader*)buf; + FuHeader* fu_header = (FuHeader*)(buf + NALU_HEADER_SIZE); + uint8_t reconstructed_nalu_type = (fu_indicator->f << 7) | + (fu_indicator->nri << 5) | + fu_header->type; + buf += NALU_HEADER_SIZE + FU_HEADER_SIZE; + size -= NALU_HEADER_SIZE + FU_HEADER_SIZE; + if (fu_header->s) { + memcpy(nalu_buf, &nalu_start_4bytecode, NALU_START_CODE_SIZE); + *nalu_offset = NALU_START_CODE_SIZE; + memcpy(nalu_buf + *nalu_offset, &reconstructed_nalu_type, 1); + *nalu_offset += 1; + memcpy(nalu_buf + *nalu_offset, buf, size); + *nalu_offset += size; + } else if (*nalu_offset < CONFIG_MAX_NALU_SIZE) { + memcpy(nalu_buf + *nalu_offset, buf, size); + *nalu_offset += size; + if (fu_header->e) { + // end of fragmented NALU + if (rtp_decoder->on_packet != NULL) { + rtp_decoder->on_packet(nalu_buf, *nalu_offset, rtp_decoder->user_data); + } + *nalu_offset = 0; // reset for next NALU + } + } + return 0; +} + static int rtp_decode_h264(RtpDecoder* rtp_decoder, uint8_t* buf, size_t size) { - static const uint32_t nalu_start_4bytecode = 0x01000000; static uint8_t nalu_buf[CONFIG_MAX_NALU_SIZE]; static int offset = 0; RtpPacket* rtp_packet = (RtpPacket*)buf; uint8_t nalu_type = *rtp_packet->payload & 0x1f; int payload_size = size - sizeof(RtpHeader); - if (nalu_type > 0 && nalu_type < 24) { - // NALU type 1-23 are single NALUs - memcpy(nalu_buf, &nalu_start_4bytecode, sizeof(nalu_start_4bytecode)); - offset = sizeof(nalu_start_4bytecode); - memcpy(nalu_buf + offset, rtp_packet->payload, payload_size); - offset += payload_size; - if (rtp_decoder->on_packet != NULL) { - rtp_decoder->on_packet(nalu_buf, offset, rtp_decoder->user_data); - } - return (int)size; - } else { - NaluHeader* fu_indicator = (NaluHeader*)rtp_packet->payload; - FuHeader* fu_header = (FuHeader*)(rtp_packet->payload + sizeof(NaluHeader)); - uint8_t reconstructed_nalu_type = (fu_indicator->f << 7) | - (fu_indicator->nri << 5) | - fu_header->type; - payload_size -= sizeof(NaluHeader) + sizeof(FuHeader); - if (fu_header->s) { - memcpy(nalu_buf, &nalu_start_4bytecode, sizeof(nalu_start_4bytecode)); - offset = sizeof(nalu_start_4bytecode); - memcpy(nalu_buf + offset, &reconstructed_nalu_type, 1); - offset += 1; - memcpy(nalu_buf + offset, rtp_packet->payload + 2, payload_size); - offset += payload_size; - } else if (offset < CONFIG_MAX_NALU_SIZE) { - memcpy(nalu_buf + offset, rtp_packet->payload + 2, payload_size); - offset += payload_size; - if (fu_header->e) { - // end of fragmented NALU - if (rtp_decoder->on_packet != NULL) { - rtp_decoder->on_packet(nalu_buf, offset, rtp_decoder->user_data); - } - offset = 0; // reset for next NALU - } - } + + switch (nalu_type) { + case STAP_A: + return rtp_decode_h264_stap_a(rtp_decoder, rtp_packet->payload + 1, payload_size - 1, nalu_buf, &offset); + case FU_A: + return rtp_decode_h264_fu_a(rtp_decoder, rtp_packet->payload, payload_size, nalu_buf, &offset); + default: + return rtp_decode_h264_single(rtp_decoder, rtp_packet->payload, payload_size, nalu_buf, &offset); + break; } return 0; } diff --git a/src/rtp.h b/src/rtp.h index 4f946145..db6a11d1 100644 --- a/src/rtp.h +++ b/src/rtp.h @@ -35,21 +35,8 @@ typedef enum RtpSsrc { } RtpSsrc; typedef struct RtpHeader { -#if __BYTE_ORDER == __BIG_ENDIAN - uint16_t version : 2; - uint16_t padding : 1; - uint16_t extension : 1; - uint16_t csrccount : 4; - uint16_t markerbit : 1; - uint16_t type : 7; -#elif __BYTE_ORDER == __LITTLE_ENDIAN - uint16_t csrccount : 4; - uint16_t extension : 1; - uint16_t padding : 1; - uint16_t version : 2; - uint16_t type : 7; - uint16_t markerbit : 1; -#endif + uint8_t vpxcc; /* Version, Padding, Extension, CSRC count */ + uint8_t mpt; /* Marker, Payload Type */ uint16_t seq_number; uint32_t timestamp; uint32_t ssrc; @@ -57,6 +44,19 @@ typedef struct RtpHeader { } RtpHeader; +static inline void rtp_header_init(RtpHeader* header, uint8_t type) { + header->vpxcc = 0x80U; + header->mpt = type & 0x7fU; +} + +static inline uint8_t rtp_header_type(const RtpHeader* header) { + return header->mpt & 0x7fU; +} + +static inline void rtp_header_set_marker(RtpHeader* header) { + header->mpt |= 0x80U; +} + typedef struct RtpPacket { RtpHeader header; uint8_t payload[0]; diff --git a/src/sctp.c b/src/sctp.c index 9819183b..b7f0acf0 100644 --- a/src/sctp.c +++ b/src/sctp.c @@ -93,6 +93,29 @@ static int sctp_outgoing_data_cb(void* userdata, void* buf, size_t len, uint8_t return 0; } +#if !CONFIG_USE_USRSCTP +static int sctp_next_stream_sequence(Sctp* sctp, uint16_t sid, uint16_t* sequence) { + int i; + + for (i = 0; i < sctp->outgoing_stream_count; i++) { + if (sctp->outgoing_streams[i].sid == sid) { + *sequence = sctp->outgoing_streams[i].next_sequence++; + return 0; + } + } + + if (sctp->outgoing_stream_count >= SCTP_MAX_STREAMS) { + return -1; + } + + sctp->outgoing_streams[sctp->outgoing_stream_count].sid = sid; + sctp->outgoing_streams[sctp->outgoing_stream_count].next_sequence = 1; + sctp->outgoing_stream_count++; + *sequence = 0; + return 0; +} +#endif + int sctp_outgoing_data(Sctp* sctp, char* buf, size_t len, SctpDataPpid ppid, uint16_t sid) { #if CONFIG_USE_USRSCTP int res; @@ -113,7 +136,11 @@ int sctp_outgoing_data(Sctp* sctp, char* buf, size_t len, SctpDataPpid ppid, uin size_t padding_len = 0; size_t payload_max = SCTP_MTU - sizeof(SctpPacket) - sizeof(SctpDataChunk); size_t pos = 0; - static uint16_t sqn = 0; + uint16_t stream_seq = 0; + + if (ppid == PPID_CONTROL && sctp_next_stream_sequence(sctp, sid, &stream_seq) != 0) { + return -1; + } SctpPacket* packet = (SctpPacket*)(sctp->buf); SctpDataChunk* chunk = (SctpDataChunk*)(packet->chunks); @@ -123,9 +150,10 @@ int sctp_outgoing_data(Sctp* sctp, char* buf, size_t len, SctpDataPpid ppid, uin packet->header.verification_tag = sctp->verification_tag; chunk->type = SCTP_DATA; - chunk->iube = 0x06; - chunk->sid = htons(0); - chunk->sqn = htons(sqn++); + /* DCEP OPEN/ACK control messages must be sent reliably and in order. */ + chunk->iube = ppid == PPID_CONTROL ? 0x02 : 0x06; + chunk->sid = htons(sid); + chunk->sqn = htons(stream_seq); chunk->ppid = htonl(ppid); while (len > payload_max) { @@ -137,7 +165,7 @@ int sctp_outgoing_data(Sctp* sctp, char* buf, size_t len, SctpDataPpid ppid, uin packet->header.checksum = sctp_get_checksum(sctp, (const uint8_t*)sctp->buf, SCTP_MTU); sctp_outgoing_data_cb(sctp, sctp->buf, SCTP_MTU, 0, 0); - chunk->iube = 0x04; + chunk->iube = ppid == PPID_CONTROL ? 0x00 : 0x04; len -= payload_max; pos += payload_max; } @@ -213,15 +241,19 @@ void sctp_handle_sctp_packet(Sctp* sctp, char* buf, size_t len) { } void sctp_incoming_data(Sctp* sctp, char* buf, size_t len) { - if (!sctp) + if (!sctp || !buf) return; #if CONFIG_USE_USRSCTP sctp_handle_sctp_packet(sctp, buf, len); usrsctp_conninput(sctp, buf, len, 0); #else + if (len < sizeof(SctpHeader)) + return; + size_t length = 0; size_t pos = sizeof(SctpHeader); + uint16_t chunk_len; SctpChunkCommon* chunk_common; SctpPacket* in_packet = (SctpPacket*)buf; SctpPacket* out_packet = (SctpPacket*)sctp->buf; @@ -244,42 +276,65 @@ void sctp_incoming_data(Sctp* sctp, char* buf, size_t len) { // prepare outgoing packet memset(sctp->buf, 0, sizeof(sctp->buf)); - while ((4 * (pos + 3) / 4) < len) { + while (pos + sizeof(SctpChunkCommon) <= len) { chunk_common = (SctpChunkCommon*)(buf + pos); + chunk_len = ntohs(chunk_common->length); + if (chunk_len < sizeof(SctpChunkCommon) || pos + chunk_len > len) { + LOGW("Invalid SCTP chunk length=%u", chunk_len); + return; + } switch (chunk_common->type) { case SCTP_DATA: { SctpDataChunk* data_chunk = (SctpDataChunk*)(buf + pos); SctpSackChunk* sack_chunk = (SctpSackChunk*)out_packet->chunks; + uint16_t payload_len; + uint32_t ppid; + uint16_t sid; + + if (chunk_len < sizeof(SctpDataChunk)) { + return; + } + payload_len = chunk_len - sizeof(SctpDataChunk); + ppid = ntohl(data_chunk->ppid); + sid = ntohs(data_chunk->sid); sack_chunk->common.type = SCTP_SACK; sack_chunk->common.flags = 0x00; sack_chunk->common.length = htons(16); sack_chunk->cumulative_tsn_ack = data_chunk->tsn; - sack_chunk->a_rwnd = htonl(0x02); + sack_chunk->a_rwnd = htonl((uint32_t)sizeof(sctp->buf)); length = ntohs(sack_chunk->common.length) + sizeof(SctpHeader); - LOGD("SCTP_DATA. ppid = %ld, data = %.2x", ntohl(data_chunk->ppid), data_chunk->data[0]); - if (ntohl(data_chunk->ppid) == DATA_CHANNEL_PPID_CONTROL && data_chunk->data[0] == DATA_CHANNEL_OPEN) { + if (ppid == DATA_CHANNEL_PPID_CONTROL && payload_len > 0 && + data_chunk->data[0] == DATA_CHANNEL_OPEN) { + uint16_t ack_sequence; + + if (sctp_next_stream_sequence(sctp, sid, &ack_sequence) != 0) { + return; + } data_chunk = (SctpDataChunk*)sack_chunk->blocks; data_chunk->type = SCTP_DATA; data_chunk->iube = 0x03; data_chunk->tsn = htonl(sctp->tsn++); - data_chunk->sid = htons(0); - data_chunk->sqn = htons(0); + data_chunk->sid = htons(sid); + data_chunk->sqn = htons(ack_sequence); data_chunk->ppid = htonl(DATA_CHANNEL_PPID_CONTROL); data_chunk->length = htons(1 + sizeof(SctpDataChunk)); data_chunk->data[0] = DATA_CHANNEL_ACK; length += ntohs(data_chunk->length); - } else if (ntohl(data_chunk->ppid) == DATA_CHANNEL_PPID_DOMSTRING) { + } else if (ppid == DATA_CHANNEL_PPID_DOMSTRING) { if (sctp->onmessage) { - sctp->onmessage((char*)data_chunk->data, ntohs(data_chunk->length) - sizeof(SctpDataChunk), - sctp->userdata, ntohs(data_chunk->sid)); + sctp->onmessage((char*)data_chunk->data, payload_len, + sctp->userdata, sid); } } pos = len; // Do not handle other msg } break; case SCTP_INIT: { + if (chunk_len < sizeof(SctpInitChunk)) { + return; + } LOGD("SCTP_INIT"); SctpInitChunk* init_chunk; @@ -303,14 +358,11 @@ void sctp_incoming_data(Sctp* sctp, char* buf, size_t len) { *(uint32_t*)¶m->value = htonl(0x02); length = ntohs(init_ack->common.length) + sizeof(SctpHeader); - if (!sctp->connected) { - sctp->connected = 1; - if (sctp->onopen) { - sctp->onopen(sctp->userdata); - } - } } break; case SCTP_INIT_ACK: { + if (chunk_len < sizeof(SctpInitChunk)) { + return; + } SctpInitChunk* init_ack = (SctpInitChunk*)in_packet->chunks; SctpCookieEchoChunk* cookie_echo = (SctpCookieEchoChunk*)out_packet->chunks; SctpChunkParam* param = NULL; @@ -336,14 +388,11 @@ void sctp_incoming_data(Sctp* sctp, char* buf, size_t len) { memcpy(cookie_echo->cookie, param->value, ntohs(param->length) - 4); length = ntohs(cookie_echo->common.length) + sizeof(SctpHeader); - if (!sctp->connected) { - sctp->connected = 1; - if (sctp->onopen) { - sctp->onopen(sctp->userdata); - } - } } break; case SCTP_SACK: + if (chunk_len < sizeof(SctpSackChunk)) { + return; + } #if 0 LOGD("SCTP_SACK"); sack = (SctpSackChunk*)in_packet->chunks; @@ -375,6 +424,16 @@ void sctp_incoming_data(Sctp* sctp, char* buf, size_t len) { } #endif break; + case SCTP_HEARTBEAT: { + SctpChunkCommon* heartbeat_ack = (SctpChunkCommon*)out_packet->chunks; + + memcpy(heartbeat_ack, chunk_common, chunk_len); + heartbeat_ack->type = SCTP_HEARTBEAT_ACK; + length = chunk_len + sizeof(SctpHeader); + pos = len; // Echo the HEARTBEAT-INFO parameter unchanged. + } break; + case SCTP_HEARTBEAT_ACK: + break; case SCTP_COOKIE_ECHO: { LOGD("SCTP_COOKIE_ECHO"); SctpChunkCommon* common = (SctpChunkCommon*)out_packet->chunks; @@ -382,8 +441,20 @@ void sctp_incoming_data(Sctp* sctp, char* buf, size_t len) { common->length = htons(4); length = ntohs(common->length) + sizeof(SctpHeader); pos = len; // Do not handle other msg + if (!sctp->connected) { + sctp->connected = 1; + if (sctp->onopen) { + sctp->onopen(sctp->userdata); + } + } } break; case SCTP_COOKIE_ACK: { + if (!sctp->connected) { + sctp->connected = 1; + if (sctp->onopen) { + sctp->onopen(sctp->userdata); + } + } break; } case SCTP_ABORT: @@ -410,7 +481,7 @@ void sctp_incoming_data(Sctp* sctp, char* buf, size_t len) { dtls_srtp_write(sctp->dtls_srtp, sctp->buf, length); // sctp_outgoing_data_cb(sctp, sctp->buf, SCTP_MTU, 0, 0); } - pos += ntohs(chunk_common->length); + pos += 4 * ((chunk_len + 3) / 4); } #endif } @@ -509,6 +580,8 @@ int sctp_create_association(Sctp* sctp, DtlsSrtp* dtls_srtp) { sctp->local_port = 5000; sctp->remote_port = 5000; sctp->tsn = 1234; + sctp->outgoing_stream_count = 0; + memset(sctp->outgoing_streams, 0, sizeof(sctp->outgoing_streams)); #if CONFIG_USE_USRSCTP int ret = -1; usrsctp_sysctl_set_sctp_ecn_enable(0); diff --git a/src/sctp.h b/src/sctp.h index 1116996a..aa6c958d 100644 --- a/src/sctp.h +++ b/src/sctp.h @@ -147,6 +147,11 @@ typedef struct { uint16_t sid; // Stream ID } SctpStreamEntry; +typedef struct { + uint16_t sid; + uint16_t next_sequence; +} SctpOutgoingStream; + typedef struct Sctp { struct socket* sock; @@ -158,6 +163,8 @@ typedef struct Sctp { DtlsSrtp* dtls_srtp; int stream_count; SctpStreamEntry stream_table[SCTP_MAX_STREAMS]; + int outgoing_stream_count; + SctpOutgoingStream outgoing_streams[SCTP_MAX_STREAMS]; /* datachannel */ void (*onmessage)(char* msg, size_t len, void* userdata, uint16_t sid); diff --git a/src/sdp.c b/src/sdp.c index b7661a4c..9d7666c3 100644 --- a/src/sdp.c +++ b/src/sdp.c @@ -35,7 +35,19 @@ void sdp_append_h264(char* sdp) { sdp_append(sdp, "a=rtpmap:96 H264/90000"); sdp_append(sdp, "a=ssrc:1 cname:webrtc-h264"); sdp_append(sdp, "a=sendrecv"); - sdp_append(sdp, "a=mid:video"); + sdp_append(sdp, "a=mid:1"); + sdp_append(sdp, "a=rtcp-mux"); +} + +void sdp_append_vp8(char* sdp) { + sdp_append(sdp, "m=video 9 UDP/TLS/RTP/SAVPF 97"); + sdp_append(sdp, "c=IN IP4 0.0.0.0"); + sdp_append(sdp, "a=rtcp-fb:97 nack"); + sdp_append(sdp, "a=rtcp-fb:97 nack pli"); + sdp_append(sdp, "a=rtpmap:97 VP8/90000"); + sdp_append(sdp, "a=ssrc:1 cname:webrtc-vp8"); + sdp_append(sdp, "a=sendrecv"); + sdp_append(sdp, "a=mid:1"); sdp_append(sdp, "a=rtcp-mux"); } @@ -45,7 +57,7 @@ void sdp_append_pcma(char* sdp) { sdp_append(sdp, "a=rtpmap:8 PCMA/8000"); sdp_append(sdp, "a=ssrc:4 cname:webrtc-pcma"); sdp_append(sdp, "a=sendrecv"); - sdp_append(sdp, "a=mid:audio"); + sdp_append(sdp, "a=mid:2"); sdp_append(sdp, "a=rtcp-mux"); } @@ -55,7 +67,7 @@ void sdp_append_pcmu(char* sdp) { sdp_append(sdp, "a=rtpmap:0 PCMU/8000"); sdp_append(sdp, "a=ssrc:5 cname:webrtc-pcmu"); sdp_append(sdp, "a=sendrecv"); - sdp_append(sdp, "a=mid:audio"); + sdp_append(sdp, "a=mid:2"); sdp_append(sdp, "a=rtcp-mux"); } @@ -65,14 +77,14 @@ void sdp_append_opus(char* sdp) { sdp_append(sdp, "a=rtpmap:111 opus/48000/2"); sdp_append(sdp, "a=ssrc:6 cname:webrtc-opus"); sdp_append(sdp, "a=sendrecv"); - sdp_append(sdp, "a=mid:audio"); + sdp_append(sdp, "a=mid:2"); sdp_append(sdp, "a=rtcp-mux"); } void sdp_append_datachannel(char* sdp) { sdp_append(sdp, "m=application 50712 UDP/DTLS/SCTP webrtc-datachannel"); sdp_append(sdp, "c=IN IP4 0.0.0.0"); - sdp_append(sdp, "a=mid:datachannel"); + sdp_append(sdp, "a=mid:0"); sdp_append(sdp, "a=sctp-port:5000"); sdp_append(sdp, "a=max-message-size:262144"); } @@ -91,16 +103,16 @@ void sdp_create(char* sdp, int b_video, int b_audio, int b_datachannel) { strcat(bundle, "a=group:BUNDLE"); - if (b_video) { - strcat(bundle, " video"); + if (b_datachannel) { + strcat(bundle, " 0"); } - if (b_audio) { - strcat(bundle, " audio"); + if (b_video) { + strcat(bundle, " 1"); } - if (b_datachannel) { - strcat(bundle, " datachannel"); + if (b_audio) { + strcat(bundle, " 2"); } sdp_append(sdp, bundle); diff --git a/src/sdp.h b/src/sdp.h index ba85d71d..6bb65292 100644 --- a/src/sdp.h +++ b/src/sdp.h @@ -12,6 +12,8 @@ void sdp_append_h264(char* sdp); +void sdp_append_vp8(char* sdp); + void sdp_append_pcma(char* sdp); void sdp_append_pcmu(char* sdp); diff --git a/src/socket.c b/src/socket.c index 6f7d8ef2..a910cc21 100644 --- a/src/socket.c +++ b/src/socket.c @@ -1,6 +1,5 @@ #include #include -#include #include "socket.h" #include "utils.h" @@ -10,17 +9,17 @@ int udp_socket_add_multicast_group(UdpSocket* udp_socket, Address* mcast_addr) { struct ip_mreq imreq = {0}; struct in_addr iaddr = {0}; - imreq.imr_interface.s_addr = INADDR_ANY; // IPV4 only imreq.imr_multiaddr.s_addr = mcast_addr->sin.sin_addr.s_addr; + imreq.imr_interface.s_addr = INADDR_ANY; if ((ret = setsockopt(udp_socket->fd, IPPROTO_IP, IP_MULTICAST_IF, &iaddr, sizeof(struct in_addr))) < 0) { LOGE("Failed to set IP_MULTICAST_IF: %d", ret); return ret; } - if ((ret = setsockopt(udp_socket->fd, IPPROTO_IP, IP_ADD_MEMBERSHIP, &imreq, sizeof(struct ip_mreq))) < 0) { - LOGE("Failed to set IP_ADD_MEMBERSHIP: %d", ret); + if ((ret = setsockopt(udp_socket->fd, IPPROTO_IP, IP_ADD_MEMBERSHIP, &imreq, sizeof(imreq))) < 0) { + LOGE("Failed to set IP_ADD_MEMBERSHIP: %d errno=%d", ret, errno); return ret; } diff --git a/src/ssl_transport.c b/src/ssl_transport.c index 9847616f..e0b58007 100644 --- a/src/ssl_transport.c +++ b/src/ssl_transport.c @@ -3,10 +3,12 @@ #include #include -#include "mbedtls/ctr_drbg.h" #include "mbedtls/debug.h" -#include "mbedtls/entropy.h" #include "mbedtls/ssl.h" +#include "mbedtls/version.h" +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 +#include "psa/crypto.h" +#endif #include #include "config.h" @@ -26,9 +28,9 @@ static int ssl_transport_mbedtls_recv_timeout(void* ctx, unsigned char* buf, siz ret = select(((TcpSocket*)ctx)->fd + 1, &read_fds, NULL, NULL, &tv); if (ret < 0) { - return -1; + return MBEDTLS_ERR_SSL_INTERNAL_ERROR; } else if (ret == 0) { - // timeout + return MBEDTLS_ERR_SSL_TIMEOUT; } else { if (FD_ISSET(((TcpSocket*)ctx)->fd, &read_fds)) { ret = tcp_socket_recv((TcpSocket*)ctx, buf, len); @@ -56,6 +58,12 @@ int ssl_transport_connect(NetworkContext_t* net_ctx, mbedtls_ctr_drbg_init(&net_ctx->ctr_drbg); mbedtls_entropy_init(&net_ctx->entropy); +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 + if (psa_crypto_init() != PSA_SUCCESS) { + return -1; + } +#endif + if ((ret = mbedtls_ctr_drbg_seed(&net_ctx->ctr_drbg, mbedtls_entropy_func, &net_ctx->entropy, (const unsigned char*)pers, strlen(pers))) != 0) { return -1; @@ -79,7 +87,9 @@ int ssl_transport_connect(NetworkContext_t* net_ctx, mbedtls_ssl_conf_ca_chain(&net_ctx->conf, &net_ctx->cacert, NULL); */ +#if MBEDTLS_VERSION_NUMBER < 0x04000000 mbedtls_ssl_conf_rng(&net_ctx->conf, mbedtls_ctr_drbg_random, &net_ctx->ctr_drbg); +#endif if ((ret = mbedtls_ssl_setup(&net_ctx->ssl, &net_ctx->conf)) != 0) { LOGE("ssl setup error: -0x%x", (unsigned int)-ret); @@ -106,9 +116,12 @@ int ssl_transport_connect(NetworkContext_t* net_ctx, LOGI("start to handshake"); while ((ret = mbedtls_ssl_handshake(&net_ctx->ssl)) != 0) { - if (ret != MBEDTLS_ERR_SSL_WANT_READ && ret != MBEDTLS_ERR_SSL_WANT_WRITE) { - LOGE("ssl handshake error: -0x%x", (unsigned int)-ret); + if (ret == MBEDTLS_ERR_SSL_WANT_READ || ret == MBEDTLS_ERR_SSL_WANT_WRITE) { + continue; } + + LOGE("ssl handshake error: -0x%x", (unsigned int)-ret); + return -1; } LOGI("handshake success"); diff --git a/src/ssl_transport.h b/src/ssl_transport.h index fe794dc0..4bcadf5d 100644 --- a/src/ssl_transport.h +++ b/src/ssl_transport.h @@ -3,8 +3,21 @@ #ifndef DISABLE_PEER_SIGNALING +#if defined(__has_include) +#if __has_include() #include +#else +#include +#endif +#if __has_include() #include +#else +#include +#endif +#else +#include +#include +#endif #include #include diff --git a/src/utils.c b/src/utils.c index 020b8acf..f9ab38f6 100644 --- a/src/utils.c +++ b/src/utils.c @@ -5,6 +5,10 @@ #include #include #include "mbedtls/md.h" +#include "mbedtls/version.h" +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 +#include "psa/crypto.h" +#endif void utils_random_string(char* s, const int len) { int i; @@ -24,23 +28,45 @@ void utils_random_string(char* s, const int len) { } void utils_get_hmac_sha1(const char* input, size_t input_len, const char* key, size_t key_len, unsigned char* output) { - mbedtls_md_context_t ctx; - mbedtls_md_type_t md_type = MBEDTLS_MD_SHA1; - mbedtls_md_init(&ctx); - mbedtls_md_setup(&ctx, mbedtls_md_info_from_type(md_type), 1); - mbedtls_md_hmac_starts(&ctx, (const unsigned char*)key, key_len); - mbedtls_md_hmac_update(&ctx, (const unsigned char*)input, input_len); - mbedtls_md_hmac_finish(&ctx, output); - mbedtls_md_free(&ctx); + memset(output, 0, 20); +#if MBEDTLS_VERSION_NUMBER >= 0x04000000 + psa_key_attributes_t attr = PSA_KEY_ATTRIBUTES_INIT; + mbedtls_svc_key_id_t key_id = MBEDTLS_SVC_KEY_ID_INIT; + size_t out_len = 0; + + psa_set_key_usage_flags(&attr, PSA_KEY_USAGE_SIGN_MESSAGE); + psa_set_key_algorithm(&attr, PSA_ALG_HMAC(PSA_ALG_SHA_1)); + psa_set_key_type(&attr, PSA_KEY_TYPE_HMAC); + psa_set_key_bits(&attr, key_len * 8); + + if (psa_crypto_init() == PSA_SUCCESS && + psa_import_key(&attr, (const uint8_t*)key, key_len, &key_id) == PSA_SUCCESS) { + if (psa_mac_compute(key_id, PSA_ALG_HMAC(PSA_ALG_SHA_1), + (const uint8_t*)input, input_len, output, 20, &out_len) != PSA_SUCCESS || + out_len != 20) { + memset(output, 0, 20); + } + psa_destroy_key(key_id); + } + psa_reset_key_attributes(&attr); +#else + const mbedtls_md_info_t* md_info = mbedtls_md_info_from_type(MBEDTLS_MD_SHA1); + if (md_info != NULL) { + mbedtls_md_context_t ctx; + mbedtls_md_init(&ctx); + if (mbedtls_md_setup(&ctx, md_info, 1) == 0) { + mbedtls_md_hmac_starts(&ctx, (const unsigned char*)key, key_len); + mbedtls_md_hmac_update(&ctx, (const unsigned char*)input, input_len); + mbedtls_md_hmac_finish(&ctx, output); + } + mbedtls_md_free(&ctx); + } +#endif } void utils_get_md5(const char* input, size_t input_len, unsigned char* output) { - mbedtls_md_context_t ctx; - mbedtls_md_type_t md_type = MBEDTLS_MD_MD5; - mbedtls_md_init(&ctx); - mbedtls_md_setup(&ctx, mbedtls_md_info_from_type(md_type), 1); - mbedtls_md_starts(&ctx); - mbedtls_md_update(&ctx, (const unsigned char*)input, input_len); - mbedtls_md_finish(&ctx, output); - mbedtls_md_free(&ctx); + const mbedtls_md_info_t* md_info = mbedtls_md_info_from_type(MBEDTLS_MD_MD5); + if (md_info != NULL) { + mbedtls_md(md_info, (const unsigned char*)input, input_len, output); + } } diff --git a/tests/test_agent.c b/tests/test_agent.c index 6cea85fa..b7e4e22f 100644 --- a/tests/test_agent.c +++ b/tests/test_agent.c @@ -20,9 +20,9 @@ int main(int argc, char* argv[]) { Agent agent; char stunserver[] = "stun:stun.l.google.com:19302"; - char turnserver[] = ""; - char username[] = ""; - char credential[] = ""; + // char turnserver[] = ""; + // char username[] = ""; + // char credential[] = ""; char description[1024]; memset(&description, 0, sizeof(description)); @@ -30,7 +30,7 @@ int main(int argc, char* argv[]) { test_gather_host(&agent); test_gather_stun(&agent, stunserver); - test_gather_turn(&agent, turnserver, username, credential); + // test_gather_turn(&agent, turnserver, username, credential); agent_get_local_description(&agent, description, sizeof(description)); printf("sdp:\n%s\n", description); diff --git a/tests/test_peer_connection.c b/tests/test_peer_connection.c index 55526259..c1106dbc 100644 --- a/tests/test_peer_connection.c +++ b/tests/test_peer_connection.c @@ -101,14 +101,14 @@ int main(int argc, char* argv[]) { int attempts = 0, datachannel_created = 0; while (attempts < MAX_CONNECTION_ATTEMPTS) { - if (!datachannel_created && peer_connection_get_state(test_user_data.offer_peer_connection) == PEER_CONNECTION_COMPLETED) { + if (!datachannel_created && peer_connection_get_state(test_user_data.offer_peer_connection) == PEER_CONNECTION_CONNECTED) { if (peer_connection_create_datachannel(test_user_data.offer_peer_connection, DATA_CHANNEL_RELIABLE, 0, 0, DATACHANNEL_NAME, "bar") == 18) { datachannel_created = 1; } } - if (peer_connection_get_state(test_user_data.offer_peer_connection) == PEER_CONNECTION_COMPLETED && - peer_connection_get_state(test_user_data.answer_peer_connection) == PEER_CONNECTION_COMPLETED && + if (peer_connection_get_state(test_user_data.offer_peer_connection) == PEER_CONNECTION_CONNECTED && + peer_connection_get_state(test_user_data.answer_peer_connection) == PEER_CONNECTION_CONNECTED && test_user_data.onmessage_offer_called == 1 && test_user_data.onmessage_answer_called == 1) { break; diff --git a/tests/test_sdp.c b/tests/test_sdp.c index d622fff9..f08b9758 100644 --- a/tests/test_sdp.c +++ b/tests/test_sdp.c @@ -3,11 +3,11 @@ #include #include "agent.h" - +#if 0 void on_agent_state_changed(AgentState state, void* user_data) { printf("Agent state changed: %d\n", state); } - +#endif int main(int argc, char* argv[]) { #if 0 Agent agent; diff --git a/third_party/mbedtls b/third_party/mbedtls index 1873d3bf..0fe989b6 160000 --- a/third_party/mbedtls +++ b/third_party/mbedtls @@ -1 +1 @@ -Subproject commit 1873d3bfc2da771672bd8e7e8f41f57e0af77f33 +Subproject commit 0fe989b6b514192783c469039edd325fd0989806 diff --git a/zephyr/CMakeLists.txt b/zephyr/CMakeLists.txt new file mode 100644 index 00000000..d26759db --- /dev/null +++ b/zephyr/CMakeLists.txt @@ -0,0 +1,39 @@ +if(CONFIG_LIBPEER) + zephyr_library_named(libpeer) + + # Align with upstream libpeer build style: compile core C files. + file(GLOB LIBPEER_CORE_SRCS ${CMAKE_CURRENT_LIST_DIR}/../src/*.c) + + include(${CMAKE_CURRENT_LIST_DIR}/../third_party/coreHTTP/httpFilePaths.cmake) + include(${CMAKE_CURRENT_LIST_DIR}/../third_party/coreMQTT/mqttFilePaths.cmake) + + set(LIBPEER_SRCS + ${LIBPEER_CORE_SRCS} + ${HTTP_SOURCES} + ${MQTT_SOURCES} + ${MQTT_SERIALIZER_SOURCES} + ) + + zephyr_library_sources(${LIBPEER_SRCS}) + + target_compile_definitions(libpeer PRIVATE + HTTP_DO_NOT_USE_CUSTOM_CONFIG + MQTT_DO_NOT_USE_CUSTOM_CONFIG + CONFIG_USE_USRSCTP=0 + ) + + zephyr_library_include_directories( + ${CMAKE_CURRENT_LIST_DIR}/../src + ${HTTP_INCLUDE_PUBLIC_DIRS} + ${MQTT_INCLUDE_PUBLIC_DIRS} + ) + + zephyr_include_directories( + ${CMAKE_CURRENT_LIST_DIR}/../src + ${HTTP_INCLUDE_PUBLIC_DIRS} + ${MQTT_INCLUDE_PUBLIC_DIRS} + ) + + zephyr_library_link_libraries(cjson) + zephyr_library_link_libraries(srtp) +endif() diff --git a/zephyr/Kconfig b/zephyr/Kconfig new file mode 100644 index 00000000..1ab637f7 --- /dev/null +++ b/zephyr/Kconfig @@ -0,0 +1,10 @@ +menu "libpeer" + +config LIBPEER + bool "Enable libpeer module" + default n + depends on CJSON && LIBSRTP + help + Enable libpeer integration for Zephyr builds. + +endmenu diff --git a/zephyr/module.yml b/zephyr/module.yml new file mode 100644 index 00000000..cbff6a1a --- /dev/null +++ b/zephyr/module.yml @@ -0,0 +1,3 @@ +build: + cmake: zephyr + kconfig: zephyr/Kconfig