summaryrefslogtreecommitdiff
path: root/src/scripting.c
diff options
context:
space:
mode:
Diffstat (limited to 'src/scripting.c')
-rw-r--r--src/scripting.c92
1 files changed, 36 insertions, 56 deletions
diff --git a/src/scripting.c b/src/scripting.c
index e6a70e546..a781e68ea 100644
--- a/src/scripting.c
+++ b/src/scripting.c
@@ -1141,42 +1141,35 @@ int redis_math_randomseed (lua_State *L) {
* EVAL and SCRIPT commands implementation
* ------------------------------------------------------------------------- */
-/* Define a lua function with the specified function name and body.
- * The function name musts be a 42 characters long string, since all the
- * functions we defined in the Lua context are in the form:
+/* Define a Lua function with the specified body.
+ * The function name will be generated in the following form:
*
* f_<hex sha1 sum>
*
- * If 'funcname' is NULL, the function name is created by the function
- * on the fly doing the SHA1 of the body, this means that passing the funcname
- * is just an optimization in case it's already at hand.
- *
- * if 'allow_dup' is true, the function can be called with a script already
- * in memory without crashing in assert(). In this case C_OK is returned.
- *
* The function increments the reference count of the 'body' object as a
* side effect of a successful call.
*
- * On success C_OK is returned, and nothing is left on the Lua stack.
- * On error C_ERR is returned and an appropriate error is set in the
- * client context. */
-int luaCreateFunction(client *c, lua_State *lua, char *funcname, robj *body, int allow_dup) {
- char fname[43];
-
- if (funcname == NULL) {
- fname[0] = 'f';
- fname[1] = '_';
- sha1hex(fname+2,body->ptr,sdslen(body->ptr));
- funcname = fname;
- }
-
- if (allow_dup) {
- sds sha = sdsnewlen(funcname+2,40);
- if (dictFind(server.lua_scripts,sha) != NULL) {
- sdsfree(sha);
- return C_OK;
- }
+ * On success a pointer to an SDS string representing the function SHA1 of the
+ * just added function is returned (and will be valid until the next call
+ * to scriptingReset() function), otherwise NULL is returned.
+ *
+ * The function handles the fact of being called with a script that already
+ * exists, and in such a case, it behaves like in the success case.
+ *
+ * If 'c' is not NULL, on error the client is informed with an appropriate
+ * error describing the nature of the problem and the Lua interpreter error. */
+sds luaCreateFunction(client *c, lua_State *lua, robj *body) {
+ char funcname[43];
+ dictEntry *de;
+
+ funcname[0] = 'f';
+ funcname[1] = '_';
+ sha1hex(funcname+2,body->ptr,sdslen(body->ptr));
+
+ sds sha = sdsnewlen(funcname+2,40);
+ if ((de = dictFind(server.lua_scripts,sha)) != NULL) {
sdsfree(sha);
+ return dictGetKey(de);
}
sds funcdef = sdsempty();
@@ -1193,29 +1186,29 @@ int luaCreateFunction(client *c, lua_State *lua, char *funcname, robj *body, int
lua_tostring(lua,-1));
}
lua_pop(lua,1);
+ sdsfree(sha);
sdsfree(funcdef);
- return C_ERR;
+ return NULL;
}
sdsfree(funcdef);
+
if (lua_pcall(lua,0,0,0)) {
if (c != NULL) {
addReplyErrorFormat(c,"Error running script (new function): %s\n",
lua_tostring(lua,-1));
}
lua_pop(lua,1);
- return C_ERR;
+ sdsfree(sha);
+ return NULL;
}
/* We also save a SHA1 -> Original script map in a dictionary
* so that we can replicate / write in the AOF all the
* EVALSHA commands as EVAL using the original script. */
- {
- int retval = dictAdd(server.lua_scripts,
- sdsnewlen(funcname+2,40),body);
- serverAssertWithInfo(c ? c : server.lua_client,NULL,retval == DICT_OK);
- incrRefCount(body);
- }
- return C_OK;
+ int retval = dictAdd(server.lua_scripts,sha,body);
+ serverAssertWithInfo(c ? c : server.lua_client,NULL,retval == DICT_OK);
+ incrRefCount(body);
+ return sha;
}
/* This is the Lua script "count" hook that we use to detect scripts timeout. */
@@ -1314,10 +1307,10 @@ void evalGenericCommand(client *c, int evalsha) {
addReply(c, shared.noscripterr);
return;
}
- if (luaCreateFunction(c,lua,funcname,c->argv[1],0) == C_ERR) {
+ if (luaCreateFunction(c,lua,c->argv[1]) == NULL) {
lua_pop(lua,1); /* remove the error handler from the stack. */
/* The error is sent to the client by luaCreateFunction()
- * itself when it returns C_ERR. */
+ * itself when it returns NULL. */
return;
}
/* Now the following is guaranteed to return non nil */
@@ -1478,22 +1471,9 @@ void scriptCommand(client *c) {
addReply(c,shared.czero);
}
} else if (c->argc == 3 && !strcasecmp(c->argv[1]->ptr,"load")) {
- char funcname[43];
- sds sha;
-
- funcname[0] = 'f';
- funcname[1] = '_';
- sha1hex(funcname+2,c->argv[2]->ptr,sdslen(c->argv[2]->ptr));
- sha = sdsnewlen(funcname+2,40);
- if (dictFind(server.lua_scripts,sha) == NULL) {
- if (luaCreateFunction(c,server.lua,funcname,c->argv[2],0)
- == C_ERR) {
- sdsfree(sha);
- return;
- }
- }
- addReplyBulkCBuffer(c,funcname+2,40);
- sdsfree(sha);
+ sds sha = luaCreateFunction(c,server.lua,c->argv[2]);
+ if (sha == NULL) return; /* The error was sent by luaCreateFunction(). */
+ addReplyBulkCBuffer(c,sha,40);
forceCommandPropagation(c,PROPAGATE_REPL|PROPAGATE_AOF);
} else if (c->argc == 2 && !strcasecmp(c->argv[1]->ptr,"kill")) {
if (server.lua_caller == NULL) {