非専門的シンギュラリティー研究所

無限に動き続けるシステムを表す方法を AI なども使って考えていきます。

C++ の調査(4)

ChatGPT で作ってもらった構文解析のプログラムを改造しようとしていたのですが、複雑で改造が難しそうなので、新しく作り直してもらうことにしました。

以下のような構文の構文解析を行い計算をするコマンドを実行する C++ のプログラムを作ってください:

構文(の全体)は以下のどれかとします:

  • 変数名 = 式
  • コマンド

式は、数値、変数、加減乗除、「- 式」(単項のマイナス)、かっこが使えるとします。
コマンドは以下のどれかとします:

  • コマンド名
  • コマンド名 変数名
  • コマンド名 数値
  • コマンド名 変数名, 数値

コマンド名はあらかじめ決まっている文字列とします。

構文解析の結果を使って以下のようなコマンドを実行するプログラムを作ってください:

で変数に定義された式を表示します。このとき式の各「部分式」に番号をつけて表示します。

で、指定した変数の指定した部分式を表示します。
変数名を省略すると、_ (アンダースコア)という名前の変数を表示します。

ですべての変数を表示します。

eval 変数名, 部分式の番号

で、指定した変数の指定した部分式を

  • 部分式が計算できるならばその結果で置き換えます。
  • 部分式が変数ならばその定義内容で置き換えます。

変数名を省略すると、_ (アンダースコア)という名前の変数に対して上記のことを行います。
print 変数名

変数名 = ?
と書くこともできるとします。
print *

* = ?
と書くこともできるとします。
_ = 式
だけ書くこともできるとします。

と入力すると以下のプログラムを得ることができました。

#include <iostream>
#include <string>
#include <sstream>
#include <map>
#include <memory>
#include <vector>
#include <cctype>
#include <stdexcept>

using namespace std;

//---------------------- AST Node Definitions -------------------------//

struct Expr {
    virtual double eval(map<string, shared_ptr<Expr>>& vars) const = 0;
    virtual string str() const = 0;
    virtual shared_ptr<Expr> clone() const = 0;
    virtual ~Expr() {}
};

using ExprPtr = shared_ptr<Expr>;

struct Number : Expr {
    double value;
    Number(double v) : value(v) {}
    double eval(map<string, ExprPtr>&) const override { return value; }
    string str() const override { return to_string(value); }
    ExprPtr clone() const override { return make_shared<Number>(*this); }
};

struct Variable : Expr {
    string name;
    Variable(string n) : name(n) {}
    double eval(map<string, ExprPtr>& vars) const override {
        if (vars.count(name)) return vars[name]->eval(vars);
        throw runtime_error("Undefined variable: " + name);
    }
    string str() const override { return name; }
    ExprPtr clone() const override { return make_shared<Variable>(*this); }
};

struct Unary : Expr {
    char op;
    ExprPtr expr;
    Unary(char o, ExprPtr e) : op(o), expr(e) {}
    double eval(map<string, ExprPtr>& vars) const override {
        double v = expr->eval(vars);
        if (op == '-') return -v;
        throw runtime_error("Unknown unary operator");
    }
    string str() const override { return string(1, op) + expr->str(); }
    ExprPtr clone() const override { return make_shared<Unary>(op, expr->clone()); }
};

struct Binary : Expr {
    char op;
    ExprPtr lhs, rhs;
    Binary(char o, ExprPtr l, ExprPtr r) : op(o), lhs(l), rhs(r) {}
    double eval(map<string, ExprPtr>& vars) const override {
        double l = lhs->eval(vars);
        double r = rhs->eval(vars);
        switch (op) {
            case '+': return l + r;
            case '-': return l - r;
            case '*': return l * r;
            case '/': return l / r;
        }
        throw runtime_error("Unknown binary operator");
    }
    string str() const override {
        return "(" + lhs->str() + " " + op + " " + rhs->str() + ")";
    }
    ExprPtr clone() const override {
        return make_shared<Binary>(op, lhs->clone(), rhs->clone());
    }
};

//---------------------- Parser -------------------------//

class Parser {
    string input;
    size_t pos;

    char peek() const { return pos < input.size() ? input[pos] : '\0'; }
    char get() { return pos < input.size() ? input[pos++] : '\0'; }

    void skipWhitespace() {
        while (isspace(peek())) get();
    }

    ExprPtr parsePrimary() {
        skipWhitespace();
        if (isdigit(peek()) || peek() == '.') {
            double val;
            stringstream ss;
            while (isdigit(peek()) || peek() == '.') ss << get();
            ss >> val;
            return make_shared<Number>(val);
        }
        else if (isalpha(peek()) || peek() == '_') {
            string name;
            while (isalnum(peek()) || peek() == '_') name += get();
            return make_shared<Variable>(name);
        }
        else if (peek() == '(') {
            get(); // '('
            auto expr = parseExpression();
            if (get() != ')') throw runtime_error("Expected ')'");
            return expr;
        }
        else if (peek() == '-') {
            get();
            return make_shared<Unary>('-', parsePrimary());
        }
        throw runtime_error("Unexpected character in primary expression");
    }

    ExprPtr parseTerm() {
        auto lhs = parsePrimary();
        while (true) {
            skipWhitespace();
            char op = peek();
            if (op == '*' || op == '/') {
                get();
                auto rhs = parsePrimary();
                lhs = make_shared<Binary>(op, lhs, rhs);
            } else break;
        }
        return lhs;
    }

    ExprPtr parseExpression() {
        auto lhs = parseTerm();
        while (true) {
            skipWhitespace();
            char op = peek();
            if (op == '+' || op == '-') {
                get();
                auto rhs = parseTerm();
                lhs = make_shared<Binary>(op, lhs, rhs);
            } else break;
        }
        return lhs;
    }

public:
    Parser(const string& in) : input(in), pos(0) {}
    ExprPtr parse() { return parseExpression(); }
};

//---------------------- REPL and Command Execution -------------------------//

map<string, ExprPtr> variables;

void printExpr(const ExprPtr& expr, int& counter, map<int, ExprPtr>& subexprs, int indent = 0) {
    subexprs[counter] = expr;
    cout << string(indent, ' ') << "[" << counter++ << "] " << expr->str() << endl;
    if (auto bin = dynamic_pointer_cast<Binary>(expr)) {
        printExpr(bin->lhs, counter, subexprs, indent + 2);
        printExpr(bin->rhs, counter, subexprs, indent + 2);
    } else if (auto un = dynamic_pointer_cast<Unary>(expr)) {
        printExpr(un->expr, counter, subexprs, indent + 2);
    }
}

void processLine(const string& line) {
    string trimmed = line;
    while (!trimmed.empty() && isspace(trimmed.back())) trimmed.pop_back();
    if (trimmed.empty()) return;

    size_t eqPos = trimmed.find('=');
    if (eqPos != string::npos && (eqPos > 0 && isalpha(trimmed[0]))) {
        string varName = trimmed.substr(0, eqPos);
        string exprStr = trimmed.substr(eqPos + 1);
        Parser p(exprStr);
        variables[varName] = p.parse();
    } else if (eqPos == string::npos && (isdigit(trimmed[0]) || trimmed[0] == '(' || trimmed[0] == '-')) {
        Parser p(trimmed);
        variables["_"] = p.parse();
    } else if (trimmed.rfind("print", 0) == 0) {
        string args = trimmed.substr(5);
        if (args == " *" || args == "*= ?") {
            for (auto& [name, expr] : variables) {
                cout << name << " = " << expr->str() << endl;
            }
        } else {
            stringstream ss(args);
            string var;
            int sub = -1;
            char comma;
            ss >> var;
            if (var == "=") var = "_";
            if (ss >> comma >> sub) {
                map<int, ExprPtr> subs;
                int counter = 0;
                printExpr(variables[var], counter, subs);
                if (subs.count(sub)) cout << subs[sub]->str() << endl;
            } else {
                map<int, ExprPtr> subs;
                int counter = 0;
                printExpr(variables[var], counter, subs);
            }
        }
    } else if (trimmed.rfind("eval", 0) == 0) {
        string args = trimmed.substr(4);
        stringstream ss(args);
        string var;
        int sub = -1;
        char comma;
        ss >> var;
        if (var.empty() || var == ",") var = "_";
        if (ss >> comma >> sub) {
            map<int, ExprPtr> subs;
            int counter = 0;
            printExpr(variables[var], counter, subs);
            if (subs.count(sub)) variables[var] = subs[sub];
        }
    } else {
        cerr << "Unknown command or syntax: " << trimmed << endl;
    }
}

int main() {
    string line;
    while (true) {
        cout << "> ";
        if (!getline(cin, line)) break;
        try {
            processLine(line);
        } catch (exception& e) {
            cerr << "Error: " << e.what() << endl;
        }
    }
    return 0;
}