mirror of
https://github.com/openai/codex.git
synced 2026-09-10 20:26:47 +00:00
fix(core): refine next-prompt suggestion requests
This commit is contained in:
@@ -20,6 +20,25 @@ use tokio::sync::mpsc;
|
||||
pub const WS_REQUEST_HEADER_TRACEPARENT_CLIENT_METADATA_KEY: &str = "ws_request_header_traceparent";
|
||||
pub const WS_REQUEST_HEADER_TRACESTATE_CLIENT_METADATA_KEY: &str = "ws_request_header_tracestate";
|
||||
|
||||
pub(crate) fn insert_max_output_tokens(
|
||||
body: &mut Value,
|
||||
max_output_tokens: Option<u64>,
|
||||
) -> Result<(), ApiError> {
|
||||
let Some(max_output_tokens) = max_output_tokens else {
|
||||
return Ok(());
|
||||
};
|
||||
let Some(body) = body.as_object_mut() else {
|
||||
return Err(ApiError::Stream(
|
||||
"failed to add max_output_tokens to responses request".to_string(),
|
||||
));
|
||||
};
|
||||
body.insert(
|
||||
"max_output_tokens".to_string(),
|
||||
Value::from(max_output_tokens),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Canonical input payload for the compaction endpoint.
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub struct CompactionInput<'a> {
|
||||
@@ -180,8 +199,6 @@ pub struct ResponsesApiRequest {
|
||||
pub stream: bool,
|
||||
pub include: Vec<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_output_tokens: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub service_tier: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub prompt_cache_key: Option<String>,
|
||||
@@ -205,7 +222,6 @@ impl From<&ResponsesApiRequest> for ResponseCreateWsRequest {
|
||||
store: request.store,
|
||||
stream: request.stream,
|
||||
include: request.include.clone(),
|
||||
max_output_tokens: request.max_output_tokens,
|
||||
service_tier: request.service_tier.clone(),
|
||||
prompt_cache_key: request.prompt_cache_key.clone(),
|
||||
text: request.text.clone(),
|
||||
@@ -231,8 +247,6 @@ pub struct ResponseCreateWsRequest {
|
||||
pub stream: bool,
|
||||
pub include: Vec<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_output_tokens: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub service_tier: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub prompt_cache_key: Option<String>,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use crate::auth::SharedAuthProvider;
|
||||
use crate::common::ResponseStream;
|
||||
use crate::common::ResponsesApiRequest;
|
||||
use crate::common::insert_max_output_tokens;
|
||||
use crate::endpoint::session::EndpointSession;
|
||||
use crate::error::ApiError;
|
||||
use crate::provider::Provider;
|
||||
@@ -71,6 +72,18 @@ impl<T: HttpTransport> ResponsesClient<T> {
|
||||
&self,
|
||||
request: ResponsesApiRequest,
|
||||
options: ResponsesOptions,
|
||||
) -> Result<ResponseStream, ApiError> {
|
||||
self.stream_request_with_max_output_tokens(
|
||||
request, options, /*max_output_tokens*/ None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn stream_request_with_max_output_tokens(
|
||||
&self,
|
||||
request: ResponsesApiRequest,
|
||||
options: ResponsesOptions,
|
||||
max_output_tokens: Option<u64>,
|
||||
) -> Result<ResponseStream, ApiError> {
|
||||
let ResponsesOptions {
|
||||
session_id,
|
||||
@@ -83,6 +96,7 @@ impl<T: HttpTransport> ResponsesClient<T> {
|
||||
|
||||
let mut body = serde_json::to_value(&request)
|
||||
.map_err(|e| ApiError::Stream(format!("failed to encode responses request: {e}")))?;
|
||||
insert_max_output_tokens(&mut body, max_output_tokens)?;
|
||||
if request.store && self.session.provider().is_azure_responses_endpoint() {
|
||||
attach_item_ids(&mut body, &request.input);
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ use crate::common::ResponseEvent;
|
||||
use crate::common::ResponseProcessedWsRequest;
|
||||
use crate::common::ResponseStream;
|
||||
use crate::common::ResponsesWsRequest;
|
||||
use crate::common::insert_max_output_tokens;
|
||||
use crate::error::ApiError;
|
||||
use crate::provider::Provider;
|
||||
use crate::rate_limits::parse_rate_limit_event;
|
||||
@@ -250,6 +251,20 @@ impl ResponsesWebsocketConnection {
|
||||
&self,
|
||||
request: ResponsesWsRequest,
|
||||
connection_reused: bool,
|
||||
) -> Result<ResponseStream, ApiError> {
|
||||
self.stream_request_with_max_output_tokens(
|
||||
request,
|
||||
connection_reused,
|
||||
/*max_output_tokens*/ None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn stream_request_with_max_output_tokens(
|
||||
&self,
|
||||
request: ResponsesWsRequest,
|
||||
connection_reused: bool,
|
||||
max_output_tokens: Option<u64>,
|
||||
) -> Result<ResponseStream, ApiError> {
|
||||
let (tx_event, rx_event) =
|
||||
mpsc::channel::<std::result::Result<ResponseEvent, ApiError>>(1600);
|
||||
@@ -259,9 +274,10 @@ impl ResponsesWebsocketConnection {
|
||||
let models_etag = self.models_etag.clone();
|
||||
let server_model = self.server_model.clone();
|
||||
let telemetry = self.telemetry.clone();
|
||||
let request_body = serde_json::to_value(&request).map_err(|err| {
|
||||
let mut request_body = serde_json::to_value(&request).map_err(|err| {
|
||||
ApiError::Stream(format!("failed to encode websocket request: {err}"))
|
||||
})?;
|
||||
insert_max_output_tokens(&mut request_body, max_output_tokens)?;
|
||||
|
||||
let current_span = Span::current();
|
||||
tokio::spawn(
|
||||
|
||||
@@ -331,7 +331,6 @@ async fn streaming_client_retries_on_transport_error() -> Result<()> {
|
||||
store: false,
|
||||
stream: true,
|
||||
include: Vec::new(),
|
||||
max_output_tokens: None,
|
||||
service_tier: None,
|
||||
prompt_cache_key: None,
|
||||
text: None,
|
||||
@@ -433,7 +432,6 @@ async fn azure_default_store_attaches_ids_and_headers() -> Result<()> {
|
||||
store: true,
|
||||
stream: true,
|
||||
include: Vec::new(),
|
||||
max_output_tokens: None,
|
||||
service_tier: None,
|
||||
prompt_cache_key: None,
|
||||
text: None,
|
||||
@@ -443,7 +441,7 @@ async fn azure_default_store_attaches_ids_and_headers() -> Result<()> {
|
||||
let mut extra_headers = HeaderMap::new();
|
||||
extra_headers.insert("x-test-header", HeaderValue::from_static("present"));
|
||||
let _stream = client
|
||||
.stream_request(
|
||||
.stream_request_with_max_output_tokens(
|
||||
request,
|
||||
ResponsesOptions {
|
||||
session_id: Some("sess_123".into()),
|
||||
@@ -453,6 +451,7 @@ async fn azure_default_store_attaches_ids_and_headers() -> Result<()> {
|
||||
compression: Compression::None,
|
||||
turn_state: None,
|
||||
},
|
||||
/*max_output_tokens*/ Some(32),
|
||||
)
|
||||
.await?;
|
||||
|
||||
@@ -496,6 +495,13 @@ async fn azure_default_store_attaches_ids_and_headers() -> Result<()> {
|
||||
.and_then(|item| item.get("id"))
|
||||
.and_then(|id| id.as_str());
|
||||
assert_eq!(input_id, Some("msg_1"));
|
||||
let max_output_tokens = req
|
||||
.body
|
||||
.as_ref()
|
||||
.and_then(RequestBody::json)
|
||||
.and_then(|body| body.get("max_output_tokens"))
|
||||
.and_then(serde_json::Value::as_u64);
|
||||
assert_eq!(max_output_tokens, Some(32));
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user