common: add bounds check in common_init_result::sampler to prevent segfault on failed model load (#21082)
* common: add bounds check in common_init_result::sampler to prevent segfault on failed model load
* Revert a308e584ca
* Add regression test
* Remove regression test for init-fail sampler check
This commit is contained in:
@@ -1243,6 +1243,9 @@ llama_context * common_init_result::context() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
common_sampler * common_init_result::sampler(llama_seq_id seq_id) {
|
common_sampler * common_init_result::sampler(llama_seq_id seq_id) {
|
||||||
|
if (seq_id < 0 || seq_id >= (int) pimpl->samplers.size()) {
|
||||||
|
return nullptr;
|
||||||
|
}
|
||||||
return pimpl->samplers[seq_id].get();
|
return pimpl->samplers[seq_id].get();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -146,19 +146,13 @@ int main(int argc, char ** argv) {
|
|||||||
|
|
||||||
ctx = llama_init->context();
|
ctx = llama_init->context();
|
||||||
model = llama_init->model();
|
model = llama_init->model();
|
||||||
|
smpl = llama_init->sampler(0);
|
||||||
|
|
||||||
if (ctx == NULL) {
|
if (ctx == NULL) {
|
||||||
LOG_ERR("%s: error: unable to create context\n", __func__);
|
LOG_ERR("%s: error: unable to create context\n", __func__);
|
||||||
return 1;
|
return 1;
|
||||||
}
|
}
|
||||||
|
|
||||||
if (model == NULL) {
|
|
||||||
LOG_ERR("%s: error: unable to load model\n", __func__);
|
|
||||||
return 1;
|
|
||||||
}
|
|
||||||
|
|
||||||
smpl = llama_init->sampler(0);
|
|
||||||
|
|
||||||
llama_memory_t mem = llama_get_memory(ctx);
|
llama_memory_t mem = llama_get_memory(ctx);
|
||||||
const llama_vocab * vocab = llama_model_get_vocab(model);
|
const llama_vocab * vocab = llama_model_get_vocab(model);
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user