Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -20,8 +20,10 @@
import static com.nvidia.nvcf.util.NvcfConstants.SCOPE_INVOKE_FUNCTION;

import com.google.common.annotations.VisibleForTesting;
import com.nvidia.boot.exceptions.NotFoundException;
import com.nvidia.boot.exceptions.UnauthorizedException;
import com.nvidia.nvcf.configuration.exceptions.InvalidInvocationException;
import com.nvidia.nvcf.persistence.function.entity.FunctionType;
import com.nvidia.nvcf.proto.ClientInvokeRequest;
import com.nvidia.nvcf.proto.ClientInvokeResponse;
import com.nvidia.nvcf.proto.ClientInvokeResponse.FunctionVersion;
Expand All @@ -46,6 +48,10 @@
@RequiredArgsConstructor
public class GrpcInvocationService extends InvocationImplBase {

private static final String MESG_LLM_FUNCTION_NOT_INVOCABLE =
"Function id '%s': LLM functions cannot be invoked through this endpoint. "
+ "Use the LLM API instead.";

private final GrpcAuthService grpcAuthService;
private final AccountService accountService;
private final FunctionInvocationValidationService functionInvocationValidationService;
Expand All @@ -64,10 +70,14 @@ public void authClientInvocation(
UUID.fromString(request.getFunctionVersionId()) : null;
var ncaId = request.hasTargetNcaId() ? request.getTargetNcaId()
: accountService.getNcaId(authentication);
var functions = lookupAndValidateAccess(authentication,
ncaId,
// LLM workers do not consume the request queues this path publishes to, so a request
// would wait out the poll window and time out. Only return versions that can serve it.
var functions = withoutLlmVersions(ncaId,
functionId,
functionVersionId);
lookupAndValidateAccess(authentication,
ncaId,
functionId,
functionVersionId));

// picking the first function version for function level info.
// all function versions will be of the same function.
Expand Down Expand Up @@ -154,4 +164,20 @@ private List<FunctionContext> lookupAndValidateAccess(
}
}

private static List<FunctionContext> withoutLlmVersions(
String ncaId,
UUID functionId,
List<FunctionContext> functions) {
var invocable = functions.stream()
.filter(context -> context.targetFunction().getFunctionType() != FunctionType.LLM)
.toList();
if (invocable.isEmpty()) {
var mesg = MESG_LLM_FUNCTION_NOT_INVOCABLE.formatted(functionId);
log.warn("Rejecting classic invocation of LLM function: ncaId={}, functionId={}",
ncaId, functionId);
throw new InvalidInvocationException(ncaId, new NotFoundException(mesg));
}
return invocable;
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -44,9 +44,11 @@
import static com.nvidia.nvcf.util.TestConstants.TEST_VERSION_ID_1;
import static com.nvidia.nvcf.util.TestConstants.TEST_VERSION_ID_2;
import static com.nvidia.nvcf.util.TestConstants.TEST_VERSION_ID_3;
import static com.nvidia.nvcf.util.TestConstants.TEST_VERSION_ID_5;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatThrownBy;

import com.nvidia.nvcf.persistence.function.entity.FunctionType;
import com.nvidia.nvcf.persistence.function.entity.RateLimitUdt;
import com.nvidia.nvcf.proto.ClientInvokeRequest;
import com.nvidia.nvcf.proto.ClientInvokeResponse;
Expand Down Expand Up @@ -1088,4 +1090,111 @@ void errorResponseIncludesNcaIdInMetadata() {
channel.shutdownNow();
}

Stream<Arguments> llmFunctionVersionArgs() {
return Stream.of(
Arguments.of(TEST_VERSION_ID_1.toString()),
Arguments.of((String) null));
}

/**
* LLM workers do not consume the request queues used by classic invocation, so the auth
* call must reject LLM functions up front instead of letting the request time out.
*/
@ParameterizedTest
@MethodSource("llmFunctionVersionArgs")
void rejectsLlmFunctionForClassicInvocation(@Nullable String versionId) {
setFunctionActive(TEST_FUNCTION_ID, TEST_VERSION_ID_1);
setFunctionType(TEST_FUNCTION_ID, TEST_VERSION_ID_1, FunctionType.LLM);
var clientAuth = MOCK_OAUTH2_TOKEN_SERVER.getJwt(TEST_CLIENT_SUBJECT,
List.of(SCOPE_INVOKE_FUNCTION), 100);
var ncaIdKey = Metadata.Key.of(TAG_NCA_ID, Metadata.ASCII_STRING_MARSHALLER);

assertThatThrownBy(() -> functionAuth(clientAuth, TEST_FUNCTION_ID.toString(), versionId))
.isInstanceOf(StatusRuntimeException.class)
.satisfies(thrown -> {
var exception = (StatusRuntimeException) thrown;
assertThat(exception.getStatus().getCode())
.isEqualTo(Status.NOT_FOUND.getCode());
assertThat(exception.getStatus().getDescription())
.contains(TEST_FUNCTION_ID.toString())
.contains("LLM functions cannot be invoked through this endpoint");
assertThat(exception.getTrailers()).isNotNull();
assertThat(exception.getTrailers().get(ncaIdKey)).isEqualTo(TEST_NCA_ID);
});
}

@Test
void rejectsLlmFunctionForClassicAdminInvocation() {
setFunctionActive(TEST_FUNCTION_ID, TEST_VERSION_ID_1);
setFunctionType(TEST_FUNCTION_ID, TEST_VERSION_ID_1, FunctionType.LLM);
var adminAuth = MOCK_OAUTH2_TOKEN_SERVER.getJwt(TEST_ADMIN_SUBJECT,
List.of(ADMIN_SCOPE_INVOKE_FUNCTION), 100);

assertThatThrownBy(() -> functionAdminAuth(adminAuth, TEST_FUNCTION_ID.toString(),
TEST_VERSION_ID_1.toString()))
.isInstanceOf(StatusRuntimeException.class)
.extracting(thrown -> ((StatusRuntimeException) thrown).getStatus().getCode())
.isEqualTo(Status.NOT_FOUND.getCode());
}

@Test
void llmFunctionWithoutAccessStillReturnsNotFound() {
// TEST_FUNCTION_ID_2 belongs to another account. Access checks must run first so the
// function type of an inaccessible function is not revealed.
setFunctionActive(TEST_FUNCTION_ID_2, TEST_VERSION_ID_2);
setFunctionType(TEST_FUNCTION_ID_2, TEST_VERSION_ID_2, FunctionType.LLM);
var clientAuth = MOCK_OAUTH2_TOKEN_SERVER.getJwt(TEST_CLIENT_SUBJECT,
List.of(SCOPE_INVOKE_FUNCTION), 100);

assertThatThrownBy(() -> functionAuth(clientAuth, TEST_FUNCTION_ID_2.toString(),
TEST_VERSION_ID_2.toString()))
.isInstanceOf(StatusRuntimeException.class)
.satisfies(thrown -> {
var status = ((StatusRuntimeException) thrown).getStatus();
assertThat(status.getCode()).isEqualTo(Status.NOT_FOUND.getCode());
assertThat(status.getDescription()).doesNotContain("LLM");
});
}

/**
* Families created before type uniformity was enforced may mix LLM and classic versions.
* A versionless classic request must still route to the classic versions.
*/
@Test
void mixedFamilyReturnsOnlyClassicVersions() {
createFunctionVersion(TEST_FUNCTION_ID, TEST_VERSION_ID_5);
setFunctionActive(TEST_FUNCTION_ID, TEST_VERSION_ID_1);
setFunctionActive(TEST_FUNCTION_ID, TEST_VERSION_ID_5);
setFunctionType(TEST_FUNCTION_ID, TEST_VERSION_ID_5, FunctionType.LLM);
var clientAuth = MOCK_OAUTH2_TOKEN_SERVER.getJwt(TEST_CLIENT_SUBJECT,
List.of(SCOPE_INVOKE_FUNCTION), 100);

var clientInvokeResponse = functionAuth(clientAuth, TEST_FUNCTION_ID.toString(), null);
assertThat(clientInvokeResponse.getFunctionVersionsList())
.extracting(FunctionVersion::getFunctionVersionId)
.containsExactly(TEST_VERSION_ID_1.toString());

// explicitly targeting the LLM version is still rejected
assertThatThrownBy(() -> functionAuth(clientAuth, TEST_FUNCTION_ID.toString(),
TEST_VERSION_ID_5.toString()))
.isInstanceOf(StatusRuntimeException.class)
.extracting(thrown -> ((StatusRuntimeException) thrown).getStatus().getCode())
.isEqualTo(Status.NOT_FOUND.getCode());
}

@Test
void allowsStreamingFunctionForClassicInvocation() {
setFunctionActive(TEST_FUNCTION_ID, TEST_VERSION_ID_1);
setFunctionType(TEST_FUNCTION_ID, TEST_VERSION_ID_1, FunctionType.STREAMING);
var clientAuth = MOCK_OAUTH2_TOKEN_SERVER.getJwt(TEST_CLIENT_SUBJECT,
List.of(SCOPE_INVOKE_FUNCTION), 100);

var clientInvokeResponse = functionAuth(
clientAuth, TEST_FUNCTION_ID.toString(), TEST_VERSION_ID_1.toString());

assertThat(clientInvokeResponse.getFunctionVersionsList())
.extracting(FunctionVersion::getFunctionVersionId)
.containsExactly(TEST_VERSION_ID_1.toString());
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -333,6 +333,12 @@ protected FunctionEntity setFunctionRateLimit(UUID functionId, UUID functionVers
return functionsRepository.save(entity);
}

protected void createFunctionVersion(UUID functionId, UUID functionVersionId) {
testService.createTestFunctionEntity(functionId, functionVersionId,
TEST_NCA_ID, TEST_FUNCTION_NAME,
FunctionStatus.DEPLOYING);
}

protected FunctionEntity setFunctionActive(UUID functionId, UUID functionVersionId) {
return setFunctionStatus(functionId, functionVersionId, FunctionStatus.ACTIVE);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,10 @@ pub const VERSION_ID_1: Uuid = uuid!("26597542-1782-4a18-aa02-504ba0598202");
pub const VERSION_ID_2: Uuid = uuid!("49331123-201a-407d-9cbc-7bb328fd0295");
pub const VERSION_ID_3: Uuid = uuid!("4117e32b-1b96-4be2-8cd2-9da4047daa05");
pub const VERSION_ID_4: Uuid = uuid!("2bf046f7-37c1-40b8-b2c2-8f8d260f7645");
#[allow(unused)]
pub const LLM_FUNCTION_ID: Uuid = uuid!("5d0c8a3e-9f4b-4e62-8a1d-3b7e2c9f0a14");
#[allow(unused)]
pub const LLM_VERSION_ID: Uuid = uuid!("e2a7b4c1-6d3f-4a85-9c02-7f1e8b5d3a96");
pub const LOCALSTACK_REGION: &str = "us-east-1";
pub const ASSETS_BUCKET: &str = "assets-bucket";
pub const RESULTS_BUCKET: &str = "results-bucket";
Expand Down Expand Up @@ -141,6 +145,7 @@ async fn mock_nvcf_api() -> ApiMockServer {
functions: vec![VERSION_ID_1, VERSION_ID_2],
has_rate_limit: false,
sync_check: false,
is_llm: false,
},
),
// FUNCTION_ID_2_RATELIMIT_SYNC has rate limit and sync check
Expand All @@ -150,6 +155,7 @@ async fn mock_nvcf_api() -> ApiMockServer {
functions: vec![VERSION_ID_3],
has_rate_limit: true,
sync_check: true,
is_llm: false,
},
),
// FUNCTION_ID_3_RATELIMIT_ASYNC has rate limit and async check
Expand All @@ -159,6 +165,16 @@ async fn mock_nvcf_api() -> ApiMockServer {
functions: vec![VERSION_ID_4],
has_rate_limit: true,
sync_check: false,
is_llm: false,
},
),
(
LLM_FUNCTION_ID,
FunctionMetadata {
functions: vec![LLM_VERSION_ID],
has_rate_limit: false,
sync_check: false,
is_llm: true,
},
),
]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,8 @@ pub struct FunctionMetadata {
pub functions: Vec<Uuid>,
pub has_rate_limit: bool,
pub sync_check: bool,
/// LLM functions are rejected by nvcf-service for classic invocation.
pub is_llm: bool,
}

#[derive(Debug)]
Expand Down Expand Up @@ -82,6 +84,15 @@ impl ApiMock {
// Include nca_id in error metadata (matching nvcf-service GrpcNcaIdException behavior)
Self::status_with_nca_id(tonic::Code::NotFound, "function not found", &client.nca_id)
})?;
if function_metadata.is_llm {
return Err(Self::status_with_nca_id(
tonic::Code::NotFound,
&format!(
"Function id '{function_id}': LLM functions cannot be invoked through this endpoint. Use the LLM API instead."
),
&client.nca_id,
));
}
let requested_function_version_id = request.function_version_id;
Ok(Response::new(ClientInvokeResponse {
function_id: function_id.into(),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,19 +16,24 @@
mod mocks;

use crate::mocks::nvcf_worker_mock::PublishMode;
use async_nats::jetstream::{context::GetStreamErrorKind, ErrorCode};
use axum::{
body::Body,
http::{self, header, HeaderValue, Method, StatusCode},
};
use futures::future;
use http_body_util::BodyExt;
use mocks::{
fixtures,
nvcf_worker_mock::{
DefaultWorkHandler, DroppableBackgroundWorker, FailedWorkHandler, Worker, WorkerProperties,
},
API_KEY, FUNCTION_ID, INSTANCE_ID, VERSION_ID_1,
API_KEY, FUNCTION_ID, INSTANCE_ID, LLM_FUNCTION_ID, LLM_VERSION_ID, VERSION_ID_1,
};
use nvcf_invocation_service::app::app;
use nvcf_invocation_service::nats::NatsService;
use nvcf_invocation_service::settings::GrpcClientConfig;
use problem_details::ProblemDetails;
use std::time::{Duration, Instant};
use tokio::{task::JoinHandle, time::timeout};
use tower::{Service, ServiceExt};
Expand Down Expand Up @@ -342,3 +347,97 @@ async fn test_delayed_worker_start() -> anyhow::Result<()> {
}
Ok(())
}

/// LLM functions are rejected by the auth call, so classic invocation must fail fast with a 404
/// on every entry point and must not publish a request that no worker will consume.
#[tokio::test]
async fn test_llm_function_rejected_without_publishing() -> anyhow::Result<()> {
let (_localstack, _nats, _mock_nvcf_api, config) = fixtures().await;
let mut app = app(config.clone(), None).await?;
let app = ServiceExt::<http::Request<Body>>::ready(&mut app).await?;

let requests = [
http::Request::builder()
.method(Method::POST)
.uri(format!(
"/v2/nvcf/pexec/functions/{LLM_FUNCTION_ID}/versions/{LLM_VERSION_ID}"
))
.header(header::AUTHORIZATION, format!("Bearer {API_KEY}"))
.body(Body::from("a body"))?,
http::Request::builder()
.method(Method::POST)
.uri(format!("/v2/nvcf/pexec/functions/{LLM_FUNCTION_ID}"))
.header(header::AUTHORIZATION, format!("Bearer {API_KEY}"))
.body(Body::from("a body"))?,
http::Request::builder()
.method(Method::POST)
.uri(format!("/v2/nvcf/exec/functions/{LLM_FUNCTION_ID}"))
.header(header::AUTHORIZATION, format!("Bearer {API_KEY}"))
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(r#"{"requestBody":{"input":"hi"}}"#))?,
// host form, as used by <function-id>.<invocation host>/v1/responses
http::Request::builder()
.method(Method::POST)
.uri("/v1/responses")
.header(
header::HOST,
format!("{LLM_FUNCTION_ID}.example.nvidia.com"),
)
.header(header::AUTHORIZATION, format!("Bearer {API_KEY}"))
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(r#"{"input":"hi"}"#))?,
];

for request in requests {
let uri = request.uri().clone();
let start = Instant::now();
let response = timeout(Duration::from_secs(10), app.call(request))
.await
.unwrap_or_else(|_| panic!("{uri}: request should not wait for a worker"))?;
assert!(
start.elapsed() < Duration::from_secs(5),
"{uri}: took {:?}",
start.elapsed()
);
assert_eq!(response.status(), StatusCode::NOT_FOUND, "{uri}");
assert_eq!(
response.headers().get(header::CONTENT_TYPE),
Some(&HeaderValue::from_static("application/problem+json")),
"{uri}"
);
let body = response.into_body().collect().await?.to_bytes();
let problem: ProblemDetails = serde_json::from_slice(&body)?;
assert_eq!(
problem.status,
Some(StatusCode::NOT_FOUND),
"{uri}: {problem:?}"
);
assert!(
problem
.detail
.as_deref()
.is_some_and(|detail| detail.contains("LLM functions cannot be invoked")),
"{uri}: {problem:?}"
);
}

// publishing creates the request stream on demand, so its absence proves nothing was queued
let nats_service = NatsService::new(
&config.nats_properties,
"http://dummy.localhost",
None,
&GrpcClientConfig::default(),
)
.await?;
let stream_name = nats_service.request_stream_name(LLM_VERSION_ID);
match nats_service.jetstream().get_stream(&stream_name).await {
Ok(_) => panic!("no request stream should exist for the rejected LLM function"),
Err(err) => match err.kind() {
GetStreamErrorKind::JetStream(js_err)
if js_err.error_code() == ErrorCode::STREAM_NOT_FOUND => {}
_ => return Err(err.into()),
},
}

Ok(())
}
Loading